Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
19 changes: 8 additions & 11 deletions server/cmd/api/api/middleware.go
Original file line number Diff line number Diff line change
Expand Up @@ -19,7 +19,7 @@ type telemetryCtxKey struct{}

type telemetryRequestCtx struct {
operationID string
code string
data oapi.BrowserApiCallEventData
}

// RecordTelemetryCode attaches code submitted with the request to its api_call
Expand All @@ -30,7 +30,8 @@ func RecordTelemetryCode(ctx context.Context, code string) {
if !ok {
return
}
tc.code = events.TruncateCaptured(code, events.CapturedFieldCap)
code = events.TruncateCaptured(code, events.CapturedFieldCap)
tc.data.Code = nonEmptyString(code)
}

// Process-wide toggle for the api_call middleware. Flipped by
Expand Down Expand Up @@ -98,15 +99,11 @@ func apiCallEventData(category oapi.TelemetryEventCategory, tc *telemetryRequest
})
return data
}
eventData := oapi.BrowserApiCallEventData{
RequestId: requestID,
OperationId: tc.operationID,
Status: status,
DurationMs: durationMs,
}
if tc.code != "" {
eventData.Code = &tc.code
}
eventData := tc.data
eventData.RequestId = requestID
eventData.OperationId = tc.operationID
eventData.Status = status
eventData.DurationMs = durationMs
data, _ := json.Marshal(eventData)
return data
}
Expand Down
55 changes: 50 additions & 5 deletions server/cmd/api/api/webmcp.go
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@ import (
"time"
"unicode/utf8"

"github.com/kernel/kernel-images/server/lib/events"
"github.com/kernel/kernel-images/server/lib/logger"
"github.com/kernel/kernel-images/server/lib/oapi"
"github.com/kernel/kernel-images/server/lib/webmcpclient"
Expand Down Expand Up @@ -71,6 +72,13 @@ func (s *ApiService) InvokeWebMCPTool(ctx context.Context, request oapi.InvokeWe
if request.Body == nil {
return oapi.InvokeWebMCPTool400JSONResponse{BadRequestErrorJSONResponse: oapi.BadRequestErrorJSONResponse{Message: "request body is required"}}, nil
}
var inputJSON []byte
var result webmcpclient.InvocationResult
var invokeErr error
defer func() {
recordWebMCPTelemetry(ctx, request.Body, inputJSON, result, invokeErr)
}()

toolRefLength := utf8.RuneCountInString(request.Body.ToolRef)
if toolRefLength < 1 || toolRefLength > 128 {
return oapi.InvokeWebMCPTool400JSONResponse{BadRequestErrorJSONResponse: oapi.BadRequestErrorJSONResponse{Message: "tool_ref must be between 1 and 128 characters"}}, nil
Expand All @@ -92,20 +100,20 @@ func (s *ApiService) InvokeWebMCPTool(ctx context.Context, request oapi.InvokeWe
invokeCtx, cancel := context.WithTimeout(ctx, timeout)
defer cancel()

result, err := s.webmcp.Invoke(invokeCtx, request.Body.ToolRef, request.Body.Input)
if err != nil {
result, invokeErr = s.webmcp.Invoke(invokeCtx, request.Body.ToolRef, request.Body.Input)
if invokeErr != nil {
switch {
case errors.Is(err, webmcpclient.ErrToolNotFound):
case errors.Is(invokeErr, webmcpclient.ErrToolNotFound):
return oapi.InvokeWebMCPTool404JSONResponse{NotFoundErrorJSONResponse: oapi.NotFoundErrorJSONResponse{Message: "WebMCP tool is no longer available; discover tools again"}}, nil
case errors.Is(err, webmcpclient.ErrOutcomeUnknown):
case errors.Is(invokeErr, webmcpclient.ErrOutcomeUnknown):
failure := oapi.WebMCPInvocationFailure{
Code: oapi.OutcomeUnknown,
Message: "the invocation started, but its final outcome could not be observed; do not retry automatically",
}
failure.InvocationId = nonEmptyString(result.InvocationID)
return oapi.InvokeWebMCPTool504JSONResponse(failure), nil
default:
logger.FromContext(ctx).Error("failed to invoke WebMCP tool", "err", err)
logger.FromContext(ctx).Error("failed to invoke WebMCP tool", "err", invokeErr)
return oapi.InvokeWebMCPTool500JSONResponse{InternalErrorJSONResponse: oapi.InternalErrorJSONResponse{Message: "failed to invoke WebMCP tool"}}, nil
}
}
Expand All @@ -124,6 +132,43 @@ func (s *ApiService) InvokeWebMCPTool(ctx context.Context, request oapi.InvokeWe
return response, nil
}

func recordWebMCPTelemetry(ctx context.Context, request *oapi.WebMCPInvokeRequest, inputJSON []byte, result webmcpclient.InvocationResult, invokeErr error) {
tc, ok := ctx.Value(telemetryCtxKey{}).(*telemetryRequestCtx)
if !ok {
return
}
captured := func(value string) *string {
return nonEmptyString(events.TruncateCaptured(value, events.CapturedFieldCap))
}
tc.data.ToolRef = captured(request.ToolRef)
tc.data.Input = captured(string(inputJSON))
tc.data.TimeoutSec = request.TimeoutSec
tc.data.ToolName = captured(result.ToolName)
tc.data.InvocationId = captured(result.InvocationID)
tc.data.ErrorText = captured(result.ErrorText)
if result.Source != nil {
tc.data.ToolSource = &oapi.BrowserWebMCPToolSource{
WindowId: result.Source.WindowID,
TabId: result.Source.TabID,
PageUrl: events.TruncateCaptured(result.Source.PageURL, events.CapturedFieldCap),
}
if frame := result.Source.Frame; frame != nil {
tc.data.ToolSource.Frame = &oapi.WebMCPToolFrame{
FrameId: frame.FrameID,
Url: events.TruncateCaptured(frame.URL, events.CapturedFieldCap),
}
}
}
status := oapi.BrowserWebMCPInvocationStatus(strings.ToLower(result.Status))
if errors.Is(invokeErr, webmcpclient.ErrOutcomeUnknown) {
status = oapi.BrowserWebMCPInvocationStatusOutcomeUnknown
tc.data.ErrorCode = nonEmptyString("outcome_unknown")
}
if status.Valid() {
tc.data.InvocationStatus = &status
}
}

func nonEmptyString(value string) *string {
if value == "" {
return nil
Expand Down
196 changes: 196 additions & 0 deletions server/cmd/api/api/webmcp_telemetry_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,196 @@
package api

import (
"context"
"encoding/json"
"errors"
"net/http"
"net/http/httptest"
"strings"
"testing"
"unicode/utf8"

"github.com/go-chi/chi/v5"
chiMiddleware "github.com/go-chi/chi/v5/middleware"
"github.com/kernel/kernel-images/server/lib/events"
"github.com/kernel/kernel-images/server/lib/oapi"
"github.com/kernel/kernel-images/server/lib/webmcpclient"
"github.com/stretchr/testify/require"
)

func webMCPTelemetryHandler(client *fakeWebMCPClient, publish func(events.Event) (events.Envelope, bool)) http.Handler {
r := chi.NewRouter()
r.Use(chiMiddleware.RequestID, TelemetryHTTPMiddleware(publish), WebMCPRequestSizeMiddleware)
strict := oapi.NewStrictHandlerWithOptions(&ApiService{webmcp: client}, []oapi.StrictMiddlewareFunc{
TelemetryStrictMiddleware(),
}, oapi.StrictHTTPServerOptions{
RequestErrorHandlerFunc: StrictRequestErrorHandler,
ResponseErrorHandlerFunc: StrictResponseErrorHandler,
})
return oapi.HandlerFromMux(strict, r)
}

func TestWebMCPTelemetryDiscoveryIsMetadataOnly(t *testing.T) {
withTelemetryMiddlewareEnabled(t)
for _, client := range []*fakeWebMCPClient{
{tools: []webmcpclient.Tool{{Ref: "wmcp_test", Name: "private_tool"}}},
{toolsErr: webmcpclient.ErrNoPageTarget},
} {
rp := &recordingPublisher{}
rec := httptest.NewRecorder()
webMCPTelemetryHandler(client, rp.publish).ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/webmcp/tools", nil))
captured := rp.snapshot()
require.Len(t, captured, 1)
require.Equal(t, "api_call", captured[0].Type)
require.Equal(t, events.Control, captured[0].Category)
require.Equal(t, oapi.KernelApi, captured[0].Source.Kind)
var data map[string]any
require.NoError(t, json.Unmarshal(captured[0].Data, &data))
require.Len(t, data, 4)
require.Equal(t, "GetWebMCPTools", data["operation_id"])
require.NotEmpty(t, data["request_id"])
require.Equal(t, float64(rec.Code), data["status"])
require.GreaterOrEqual(t, data["duration_ms"].(float64), 0.0)
}
}

func TestWebMCPTelemetryInvocationOutcomes(t *testing.T) {
withTelemetryMiddlewareEnabled(t)
for _, test := range []struct {
name string
result webmcpclient.InvocationResult
err error
httpStatus int
status string
errorCode string
}{
{name: "completed", result: webmcpclient.InvocationResult{Status: "Completed", InvocationID: "inv_1"}, httpStatus: 200, status: "completed"},
{name: "canceled", result: webmcpclient.InvocationResult{Status: "Canceled", InvocationID: "inv_1"}, httpStatus: 200, status: "canceled"},
{name: "tool error", result: webmcpclient.InvocationResult{Status: "Error", InvocationID: "inv_1", ErrorText: "invalid quantity"}, httpStatus: 200, status: "error"},
{name: "form populated", result: webmcpclient.InvocationResult{Status: "awaiting_submission", InvocationID: "inv_1"}, httpStatus: 200, status: "awaiting_submission"},
{name: "unknown with id", result: webmcpclient.InvocationResult{InvocationID: "inv_1"}, err: webmcpclient.ErrOutcomeUnknown, httpStatus: 504, status: "outcome_unknown", errorCode: "outcome_unknown"},
{name: "unknown without id", err: webmcpclient.ErrOutcomeUnknown, httpStatus: 504, status: "outcome_unknown", errorCode: "outcome_unknown"},
{name: "stale reference", err: webmcpclient.ErrToolNotFound, httpStatus: 404},
{name: "transport error", err: errors.New("private transport detail"), httpStatus: 500},
{name: "invalid status", result: webmcpclient.InvocationResult{Status: "Unexpected", InvocationID: "inv_1"}, httpStatus: 500},
} {
t.Run(test.name, func(t *testing.T) {
rp := &recordingPublisher{}
client := &fakeWebMCPClient{result: test.result, invokeErr: test.err}
rec := httptest.NewRecorder()
webMCPTelemetryHandler(client, rp.publish).ServeHTTP(rec, httptest.NewRequest(http.MethodPost, "/webmcp/invoke", strings.NewReader(`{"tool_ref":"wmcp_test","input":{"quantity":2},"timeout_sec":30}`)))
require.Equal(t, test.httpStatus, rec.Code)
captured := rp.snapshot()
require.Len(t, captured, 1)
require.Equal(t, "api_call", captured[0].Type)
require.Equal(t, events.Control, captured[0].Category)
var data oapi.BrowserApiCallEventData
require.NoError(t, json.Unmarshal(captured[0].Data, &data))
require.Equal(t, "InvokeWebMCPTool", data.OperationId)
require.Equal(t, test.httpStatus, data.Status)
require.Equal(t, "wmcp_test", *data.ToolRef)
require.JSONEq(t, `{"quantity":2}`, *data.Input)
require.Equal(t, 30, *data.TimeoutSec)
require.Equal(t, nonEmptyString(test.result.InvocationID), data.InvocationId)
require.Equal(t, nonEmptyString(test.result.ErrorText), data.ErrorText)
require.Equal(t, nonEmptyString(test.errorCode), data.ErrorCode)
if test.status == "" {
require.Nil(t, data.InvocationStatus)
} else {
require.Equal(t, test.status, string(*data.InvocationStatus))
}
require.Nil(t, data.Code)
require.Nil(t, data.ToolSource)
require.Nil(t, data.ToolName)
require.NotContains(t, string(captured[0].Data), "private transport detail")
})
}
}

func TestWebMCPTelemetryClipsInputAndSource(t *testing.T) {
withTelemetryMiddlewareEnabled(t)
oversized := strings.Repeat("é", events.CapturedFieldCap)
client := &fakeWebMCPClient{result: webmcpclient.InvocationResult{
ToolName: oversized,
Source: &webmcpclient.ToolSource{
WindowID: 2, TabID: 3, PageURL: oversized,
Frame: &webmcpclient.ToolFrame{FrameID: 7, URL: oversized},
},
InvocationID: oversized, Status: "Error", ErrorText: oversized,
Output: map[string]any{"private_output": "not captured"},
}}
body, err := json.Marshal(oapi.WebMCPInvokeRequest{ToolRef: "wmcp_test", Input: map[string]any{"value": oversized}})
require.NoError(t, err)
rp := &recordingPublisher{}
webMCPTelemetryHandler(client, rp.publish).ServeHTTP(httptest.NewRecorder(), httptest.NewRequest(http.MethodPost, "/webmcp/invoke", strings.NewReader(string(body))))
captured := rp.snapshot()
require.Len(t, captured, 1)
var data oapi.BrowserApiCallEventData
require.NoError(t, json.Unmarshal(captured[0].Data, &data))
require.NotNil(t, data.ToolSource)
require.Equal(t, 2, data.ToolSource.WindowId)
require.Equal(t, 3, data.ToolSource.TabId)
require.Equal(t, 7, data.ToolSource.Frame.FrameId)
for _, value := range []string{*data.Input, *data.ToolName, *data.InvocationId, *data.ErrorText, data.ToolSource.PageUrl, data.ToolSource.Frame.Url} {
require.LessOrEqual(t, len(value), events.CapturedFieldCap)
require.True(t, utf8.ValidString(value))
require.True(t, strings.HasSuffix(value, events.TruncatedSuffix))
}
require.Nil(t, data.TimeoutSec)
require.NotContains(t, string(captured[0].Data), "private_output")
require.Equal(t, oversized, client.input["value"], "telemetry clipping must not modify the invocation")
}

func TestWebMCPTelemetryValidationFailure(t *testing.T) {
withTelemetryMiddlewareEnabled(t)
rp := &recordingPublisher{}
rec := httptest.NewRecorder()
webMCPTelemetryHandler(&fakeWebMCPClient{}, rp.publish).ServeHTTP(rec, httptest.NewRequest(http.MethodPost, "/webmcp/invoke", strings.NewReader(`{"tool_ref":"wmcp_test","input":{},"timeout_sec":0}`)))
require.Equal(t, http.StatusBadRequest, rec.Code)
captured := rp.snapshot()
require.Len(t, captured, 1)
var data oapi.BrowserApiCallEventData
require.NoError(t, json.Unmarshal(captured[0].Data, &data))
require.Equal(t, 400, data.Status)
require.Equal(t, 0, *data.TimeoutSec)
require.Equal(t, "{}", *data.Input)
require.Nil(t, data.InvocationId)
require.Nil(t, data.InvocationStatus)
}

func TestWebMCPTelemetryDisabled(t *testing.T) {
withTelemetryMiddlewareEnabled(t)
DisableTelemetryMiddleware()
rp := &recordingPublisher{}
handler := webMCPTelemetryHandler(&fakeWebMCPClient{result: webmcpclient.InvocationResult{Status: "Completed"}}, rp.publish)
handler.ServeHTTP(httptest.NewRecorder(), httptest.NewRequest(http.MethodGet, "/webmcp/tools", nil))
handler.ServeHTTP(httptest.NewRecorder(), httptest.NewRequest(http.MethodPost, "/webmcp/invoke", strings.NewReader(`{"tool_ref":"wmcp_test","input":{}}`)))
require.Empty(t, rp.snapshot())
}

func TestWebMCPTelemetryNestedExecuteEmitsOncePerRequest(t *testing.T) {
withTelemetryMiddlewareEnabled(t)
rp := &recordingPublisher{}
server := httptest.NewServer(webMCPTelemetryHandler(&fakeWebMCPClient{result: webmcpclient.InvocationResult{Status: "Completed"}}, rp.publish))
defer server.Close()
code := `await webmcp.invokeTool("wmcp_test", {});`
outer := chiHandler(t, rp.publish, "ExecutePlaywrightCode", http.StatusOK, func(ctx context.Context) {
RecordTelemetryCode(ctx, code)
resp, err := http.Post(server.URL+"/webmcp/invoke", "application/json", strings.NewReader(`{"tool_ref":"wmcp_test","input":{}}`))
require.NoError(t, err)
defer resp.Body.Close()
require.Equal(t, http.StatusOK, resp.StatusCode)
})
outer.ServeHTTP(httptest.NewRecorder(), httptest.NewRequest(http.MethodPost, "/playwright/execute", nil))
captured := rp.snapshot()
require.Len(t, captured, 2)
var innerData, outerData oapi.BrowserApiCallEventData
require.NoError(t, json.Unmarshal(captured[0].Data, &innerData))
require.NoError(t, json.Unmarshal(captured[1].Data, &outerData))
require.Equal(t, "InvokeWebMCPTool", innerData.OperationId)
require.Equal(t, "ExecutePlaywrightCode", outerData.OperationId)
require.NotEqual(t, innerData.RequestId, outerData.RequestId)
require.Equal(t, code, *outerData.Code)
require.Nil(t, outerData.ToolRef)
require.Nil(t, innerData.Code)
}
Loading
Loading