Skip to content
Merged
Show file tree
Hide file tree
Changes from 2 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
11 changes: 11 additions & 0 deletions server/cmd/api/api/api.go
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,7 @@ import (
"github.com/kernel/kernel-images/server/lib/recorder"
"github.com/kernel/kernel-images/server/lib/scaletozero"
"github.com/kernel/kernel-images/server/lib/telemetry"
"github.com/kernel/kernel-images/server/lib/webmcpclient"
)

type cdpMonitorController interface {
Expand All @@ -40,6 +41,12 @@ type OTLPExporter interface {

var _ OTLPExporter = (*events.OTLPExportController)(nil)

type webMCPClient interface {
Tools(ctx context.Context, targetID string) ([]webmcpclient.Tool, error)
Invoke(ctx context.Context, toolRef string, input map[string]any) (webmcpclient.InvocationResult, error)
Close() error
}

type ApiService struct {
// defaultRecorderID is used whenever the caller doesn't specify an explicit ID.
defaultRecorderID string
Expand Down Expand Up @@ -77,6 +84,8 @@ type ApiService struct {
// playwrightDaemonCmd holds the daemon process for cleanup
playwrightDaemonCmd *exec.Cmd

webmcp webMCPClient

// policy management
policy *policy.Policy

Expand Down Expand Up @@ -159,6 +168,7 @@ func New(
telemetrySession: telemetrySession,
cdpMonitor: mon,
otlpExport: otlpExport,
webmcp: webmcpclient.NewManager(upstreamMgr),
lifecycleCtx: ctx,
lifecycleCancel: cancel,
}, nil
Expand Down Expand Up @@ -421,6 +431,7 @@ func (s *ApiService) ListRecorders(ctx context.Context, _ oapi.ListRecordersRequ
}

func (s *ApiService) Shutdown(ctx context.Context) error {
_ = s.webmcp.Close()
s.monitorMu.Lock()
s.lifecycleCancel()
s.cdpMonitor.Stop()
Expand Down
123 changes: 123 additions & 0 deletions server/cmd/api/api/webmcp.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,123 @@
package api

import (
"context"
"encoding/json"
"errors"
"strings"
"time"

"github.com/kernel/kernel-images/server/lib/logger"
"github.com/kernel/kernel-images/server/lib/oapi"
"github.com/kernel/kernel-images/server/lib/webmcpclient"
)

const (
defaultWebMCPInvocationTimeout = 60 * time.Second
maxWebMCPInvocationTimeoutSec = 120
maxWebMCPInputBytes = 1 << 20
)

func (s *ApiService) GetWebMCPTools(ctx context.Context, request oapi.GetWebMCPToolsRequestObject) (oapi.GetWebMCPToolsResponseObject, error) {
targetID := ""
if request.Params.TargetId != nil {
targetID = *request.Params.TargetId
}
tools, err := s.webmcp.Tools(ctx, targetID)
if err != nil {
if errors.Is(err, webmcpclient.ErrNoPageTarget) {
return oapi.GetWebMCPTools404JSONResponse{NotFoundErrorJSONResponse: oapi.NotFoundErrorJSONResponse{Message: err.Error()}}, nil
}
logger.FromContext(ctx).Error("failed to discover WebMCP tools", "err", err)
return oapi.GetWebMCPTools500JSONResponse{InternalErrorJSONResponse: oapi.InternalErrorJSONResponse{Message: "failed to discover WebMCP tools"}}, nil
}

responseTools := make([]oapi.WebMCPTool, 0, len(tools))
for _, tool := range tools {
inputSchema := tool.InputSchema
if inputSchema == nil {
inputSchema = make(map[string]any)
}
responseTool := oapi.WebMCPTool{
ToolRef: tool.Ref,
Name: tool.Name,
Description: tool.Description,
InputSchema: inputSchema,
PageTargetId: tool.PageTargetID,
TargetId: tool.TargetID,
TargetType: tool.TargetType,
FrameId: tool.FrameID,
}
if tool.Annotations != nil {
responseTool.Annotations = &oapi.WebMCPToolAnnotations{
ReadOnly: tool.Annotations.ReadOnly,
UntrustedContent: tool.Annotations.UntrustedContent,
Consequential: tool.Annotations.Consequential,
Autosubmit: tool.Annotations.Autosubmit,
}
}
responseTool.TargetUrl = nonEmptyString(tool.TargetURL)
responseTool.FrameUrl = nonEmptyString(tool.FrameURL)
responseTool.ParentFrameId = nonEmptyString(tool.ParentFrameID)
responseTool.DocumentRef = nonEmptyString(tool.DocumentRef)
responseTools = append(responseTools, responseTool)
}
return oapi.GetWebMCPTools200JSONResponse{Tools: responseTools}, nil
}

func (s *ApiService) InvokeWebMCPTool(ctx context.Context, request oapi.InvokeWebMCPToolRequestObject) (oapi.InvokeWebMCPToolResponseObject, error) {
if request.Body == nil {
return oapi.InvokeWebMCPTool400JSONResponse{BadRequestErrorJSONResponse: oapi.BadRequestErrorJSONResponse{Message: "request body is required"}}, nil
}
inputJSON, err := json.Marshal(request.Body.Input)
if err != nil || len(inputJSON) > maxWebMCPInputBytes {
return oapi.InvokeWebMCPTool400JSONResponse{BadRequestErrorJSONResponse: oapi.BadRequestErrorJSONResponse{Message: "input must be valid JSON no larger than 1 MiB"}}, nil
}
timeout := defaultWebMCPInvocationTimeout
if request.Body.TimeoutSec != nil {
if *request.Body.TimeoutSec < 1 || *request.Body.TimeoutSec > maxWebMCPInvocationTimeoutSec {
return oapi.InvokeWebMCPTool400JSONResponse{BadRequestErrorJSONResponse: oapi.BadRequestErrorJSONResponse{Message: "timeout_sec must be between 1 and 120"}}, nil
}
timeout = time.Duration(*request.Body.TimeoutSec) * time.Second
}
invokeCtx, cancel := context.WithTimeout(ctx, timeout)
defer cancel()

result, err := s.webmcp.Invoke(invokeCtx, request.Body.ToolRef, request.Body.Input)
if err != nil {
switch {
case errors.Is(err, 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):
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)
return oapi.InvokeWebMCPTool500JSONResponse{InternalErrorJSONResponse: oapi.InternalErrorJSONResponse{Message: "failed to invoke WebMCP tool"}}, nil
}
}

status := oapi.WebMCPInvocationResultStatus(strings.ToLower(result.Status))
if !status.Valid() {
logger.FromContext(ctx).Error("WebMCP tool returned unknown status", "status", result.Status)
return oapi.InvokeWebMCPTool500JSONResponse{InternalErrorJSONResponse: oapi.InternalErrorJSONResponse{Message: "WebMCP tool returned an unknown status"}}, nil
}
response := oapi.InvokeWebMCPTool200JSONResponse{
InvocationId: result.InvocationID,
Status: status,
Output: result.Output,
}
response.ErrorText = nonEmptyString(result.ErrorText)
return response, nil
}

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

import (
"context"
"errors"
"fmt"
"strings"
"testing"

"github.com/kernel/kernel-images/server/lib/oapi"
"github.com/kernel/kernel-images/server/lib/webmcpclient"
"github.com/stretchr/testify/require"
)

type fakeWebMCPClient struct {
tools []webmcpclient.Tool
toolsErr error
result webmcpclient.InvocationResult
invokeErr error
targetID string
toolRef string
input map[string]any
}

func (f *fakeWebMCPClient) Tools(_ context.Context, targetID string) ([]webmcpclient.Tool, error) {
f.targetID = targetID
return f.tools, f.toolsErr
}

func (f *fakeWebMCPClient) Invoke(_ context.Context, toolRef string, input map[string]any) (webmcpclient.InvocationResult, error) {
f.toolRef = toolRef
f.input = input
return f.result, f.invokeErr
}

func (f *fakeWebMCPClient) Close() error { return nil }

func TestGetWebMCPToolsMapsRegistrationContext(t *testing.T) {
client := &fakeWebMCPClient{tools: []webmcpclient.Tool{{
Ref: "wmcp_test",
Name: "pay",
Description: "Pay for the order",
Annotations: &webmcpclient.Annotations{Consequential: true},
PageTargetID: "page-target",
TargetID: "iframe-target",
TargetType: "iframe",
TargetURL: "https://payments.example/element",
FrameID: "iframe-frame",
FrameURL: "https://payments.example/element",
ParentFrameID: "page-frame",
DocumentRef: "iframe-target:loader-1",
}}}
service := &ApiService{webmcp: client}
targetID := "page-target"

response, err := service.GetWebMCPTools(context.Background(), oapi.GetWebMCPToolsRequestObject{
Params: oapi.GetWebMCPToolsParams{TargetId: &targetID},
})
require.NoError(t, err)
body := response.(oapi.GetWebMCPTools200JSONResponse)
require.Equal(t, targetID, client.targetID)
require.Len(t, body.Tools, 1)
tool := body.Tools[0]
require.Equal(t, "wmcp_test", tool.ToolRef)
require.Equal(t, "iframe-target", tool.TargetId)
require.Equal(t, "iframe-frame", tool.FrameId)
require.Equal(t, "iframe-target:loader-1", *tool.DocumentRef)
require.Empty(t, tool.InputSchema)
require.True(t, tool.Annotations.Consequential)
}

func TestInvokeWebMCPToolReturnsPageResult(t *testing.T) {
client := &fakeWebMCPClient{result: webmcpclient.InvocationResult{
InvocationID: "invocation-1",
Status: "Completed",
Output: map[string]any{"ok": true},
}}
service := &ApiService{webmcp: client}

response, err := service.InvokeWebMCPTool(context.Background(), oapi.InvokeWebMCPToolRequestObject{
Body: &oapi.WebMCPInvokeRequest{ToolRef: "wmcp_test", Input: map[string]any{"amount": 2900}},
})
require.NoError(t, err)
body := response.(oapi.InvokeWebMCPTool200JSONResponse)
require.Equal(t, "wmcp_test", client.toolRef)
require.Equal(t, 2900, client.input["amount"])
require.Equal(t, oapi.WebMCPInvocationResultStatusCompleted, body.Status)
require.Equal(t, true, body.Output.(map[string]any)["ok"])
}

func TestInvokeWebMCPToolReportsUnknownOutcome(t *testing.T) {
client := &fakeWebMCPClient{
result: webmcpclient.InvocationResult{InvocationID: "invocation-1"},
invokeErr: webmcpclient.ErrOutcomeUnknown,
}
service := &ApiService{webmcp: client}

response, err := service.InvokeWebMCPTool(context.Background(), oapi.InvokeWebMCPToolRequestObject{
Body: &oapi.WebMCPInvokeRequest{ToolRef: "wmcp_test", Input: map[string]any{}},
})
require.NoError(t, err)
body := response.(oapi.InvokeWebMCPTool504JSONResponse)
require.Equal(t, oapi.OutcomeUnknown, body.Code)
require.Equal(t, "invocation-1", *body.InvocationId)
}

func TestGetWebMCPToolsReturnsNotFoundWithoutPage(t *testing.T) {
service := &ApiService{webmcp: &fakeWebMCPClient{toolsErr: webmcpclient.ErrNoPageTarget}}
response, err := service.GetWebMCPTools(context.Background(), oapi.GetWebMCPToolsRequestObject{})
require.NoError(t, err)
_, ok := response.(oapi.GetWebMCPTools404JSONResponse)
require.True(t, ok)
}

func TestInvokeWebMCPToolReturnsNotFoundForStaleReference(t *testing.T) {
service := &ApiService{webmcp: &fakeWebMCPClient{invokeErr: webmcpclient.ErrToolNotFound}}
response, err := service.InvokeWebMCPTool(context.Background(), oapi.InvokeWebMCPToolRequestObject{
Body: &oapi.WebMCPInvokeRequest{ToolRef: "wmcp_stale", Input: map[string]any{}},
})
require.NoError(t, err)
_, ok := response.(oapi.InvokeWebMCPTool404JSONResponse)
require.True(t, ok)
}

func TestInvokeWebMCPToolRejectsTimeoutOutsideBounds(t *testing.T) {
for _, timeoutSec := range []int{0, -1, 121} {
t.Run(fmt.Sprintf("timeout_%d", timeoutSec), func(t *testing.T) {
client := &fakeWebMCPClient{}
service := &ApiService{webmcp: client}
response, err := service.InvokeWebMCPTool(context.Background(), oapi.InvokeWebMCPToolRequestObject{
Body: &oapi.WebMCPInvokeRequest{ToolRef: "wmcp_test", Input: map[string]any{}, TimeoutSec: &timeoutSec},
})
require.NoError(t, err)
_, ok := response.(oapi.InvokeWebMCPTool400JSONResponse)
require.True(t, ok)
require.Empty(t, client.toolRef)
})
}
}

func TestInvokeWebMCPToolRejectsOversizedInput(t *testing.T) {
client := &fakeWebMCPClient{}
service := &ApiService{webmcp: client}
response, err := service.InvokeWebMCPTool(context.Background(), oapi.InvokeWebMCPToolRequestObject{
Body: &oapi.WebMCPInvokeRequest{
ToolRef: "wmcp_test",
Input: map[string]any{"value": strings.Repeat("a", maxWebMCPInputBytes)},
},
})
require.NoError(t, err)
_, ok := response.(oapi.InvokeWebMCPTool400JSONResponse)
require.True(t, ok)
require.Empty(t, client.toolRef)
}

func TestInvokeWebMCPToolRejectsUnexpectedClientError(t *testing.T) {
service := &ApiService{webmcp: &fakeWebMCPClient{invokeErr: errors.New("CDP failed")}}
response, err := service.InvokeWebMCPTool(context.Background(), oapi.InvokeWebMCPToolRequestObject{
Body: &oapi.WebMCPInvokeRequest{ToolRef: "wmcp_test", Input: map[string]any{}},
})
require.NoError(t, err)
_, ok := response.(oapi.InvokeWebMCPTool500JSONResponse)
require.True(t, ok)
}
2 changes: 2 additions & 0 deletions server/lib/events/category_gen.go

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

Loading
Loading