mcp修复
This commit is contained in:
1 parent
801c0af692
commit
3080b17d02
9 files changed
+378
-73
No files matched your search
+11
-18
@@ -23,7 +23,7 @@ import (
|
||||
const (
|
||||
mcpProtocolVersion = "2025-06-18"
|
||||
mcpServerName = "fire-safety-ymd-spatial-readonly"
|
||||
mcpServerVersion = "0.2.0"
|
||||
mcpServerVersion = "0.2.2"
|
||||
maximumMCPBody = int64(1024 * 1024)
|
||||
maximumToolTimeout = 30 * time.Second
|
||||
|
||||
@@ -118,12 +118,6 @@ func (h *MCPHandler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
|
||||
h.logResult(requestID, "transport", "unsupported_media_type", started)
|
||||
return
|
||||
}
|
||||
if version := strings.TrimSpace(r.Header.Get("MCP-Protocol-Version")); version != "" && version != mcpProtocolVersion {
|
||||
h.writeRPCError(w, http.StatusBadRequest, nil, -32602, "MCP_PROTOCOL_VERSION_UNSUPPORTED", "Unsupported MCP protocol version.")
|
||||
h.logResult(requestID, "transport", "unsupported_protocol", started)
|
||||
return
|
||||
}
|
||||
|
||||
body, tooLarge, err := readBoundedBody(r.Body, h.maxBodyBytes)
|
||||
if err != nil {
|
||||
h.writeRPCError(w, http.StatusBadRequest, nil, -32700, "MCP_REQUEST_INVALID", "Unable to read MCP request.")
|
||||
@@ -150,11 +144,7 @@ func (h *MCPHandler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
|
||||
|
||||
switch request.Method {
|
||||
case "initialize":
|
||||
if h.handleInitialize(w, request) {
|
||||
h.logResult(requestID, "initialize", "success", started)
|
||||
} else {
|
||||
h.logResult(requestID, "initialize", "unsupported_protocol", started)
|
||||
}
|
||||
h.logResult(requestID, "initialize", h.handleInitialize(w, request), started)
|
||||
case "notifications/initialized":
|
||||
w.WriteHeader(http.StatusAccepted)
|
||||
h.logResult(requestID, "notifications/initialized", "accepted", started)
|
||||
@@ -169,13 +159,16 @@ func (h *MCPHandler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
}
|
||||
|
||||
func (h *MCPHandler) handleInitialize(w http.ResponseWriter, request rpcRequest) bool {
|
||||
func (h *MCPHandler) handleInitialize(w http.ResponseWriter, request rpcRequest) string {
|
||||
var params struct {
|
||||
ProtocolVersion string `json:"protocolVersion"`
|
||||
ProtocolVersion json.RawMessage `json:"protocolVersion"`
|
||||
}
|
||||
if len(request.Params) == 0 || json.Unmarshal(request.Params, ¶ms) != nil || params.ProtocolVersion != mcpProtocolVersion {
|
||||
h.writeRPCError(w, http.StatusOK, request.ID, -32602, "MCP_PROTOCOL_VERSION_UNSUPPORTED", "Client must support MCP protocol version 2025-06-18.")
|
||||
return false
|
||||
result := "compatibility_success"
|
||||
if len(request.Params) != 0 && json.Unmarshal(request.Params, ¶ms) == nil {
|
||||
var protocolVersion string
|
||||
if json.Unmarshal(params.ProtocolVersion, &protocolVersion) == nil && protocolVersion == mcpProtocolVersion {
|
||||
result = "direct_success"
|
||||
}
|
||||
}
|
||||
h.writeRPCResult(w, request.ID, map[string]any{
|
||||
"protocolVersion": mcpProtocolVersion,
|
||||
@@ -189,7 +182,7 @@ func (h *MCPHandler) handleInitialize(w http.ResponseWriter, request rpcRequest)
|
||||
},
|
||||
"instructions": "Read-only planning support. Source records with invalid geometries are excluded, so results may be incomplete. Place-name matches are candidates that require user confirmation before coordinate-based analysis. Treat all records as potentially stale, verify resource availability and field safety, and never present access-line candidates as confirmed routes or responsibility records as live team locations.",
|
||||
})
|
||||
return true
|
||||
return result
|
||||
}
|
||||
|
||||
func (h *MCPHandler) handleToolCall(w http.ResponseWriter, r *http.Request, request rpcRequest, requestID string, started time.Time) {
|
||||
|
||||
+244
-16
@@ -82,15 +82,6 @@ func TestMCPTransportGuards(t *testing.T) {
|
||||
wantStatus: http.StatusUnsupportedMediaType,
|
||||
wantErrorCode: "MCP_CONTENT_TYPE_INVALID",
|
||||
},
|
||||
{
|
||||
name: "unsupported protocol header",
|
||||
body: `{"jsonrpc":"2.0","id":1,"method":"tools/list"}`,
|
||||
configure: func(request *http.Request) {
|
||||
request.Header.Set("MCP-Protocol-Version", "1999-01-01")
|
||||
},
|
||||
wantStatus: http.StatusBadRequest,
|
||||
wantErrorCode: "MCP_PROTOCOL_VERSION_UNSUPPORTED",
|
||||
},
|
||||
{
|
||||
name: "invalid JSON",
|
||||
body: `{"jsonrpc":`,
|
||||
@@ -174,15 +165,252 @@ func TestMCPInitializeAndInitializedNotification(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestMCPRejectsUnsupportedInitializeVersion(t *testing.T) {
|
||||
handler := newTestMCPHandler(t, &fakeMCPSpatialService{}, 4096, time.Second)
|
||||
request := authenticatedMCPRequest(`{"jsonrpc":"2.0","id":1,"method":"initialize","params":{"protocolVersion":"2025-11-25"}}`)
|
||||
response := httptest.NewRecorder()
|
||||
func TestMCPProvenProfileCompletesToolCallLifecycle(t *testing.T) {
|
||||
spatial := &fakeMCPSpatialService{}
|
||||
handler := newTestMCPHandler(t, spatial, 4096, time.Second)
|
||||
tests := []struct {
|
||||
name string
|
||||
body string
|
||||
wantStatus int
|
||||
wantBody string
|
||||
}{
|
||||
{
|
||||
name: "initialize without configurable protocol version",
|
||||
body: `{"jsonrpc":"2.0","id":1,"method":"initialize"}`,
|
||||
wantStatus: http.StatusOK,
|
||||
wantBody: `"protocolVersion":"2025-06-18"`,
|
||||
},
|
||||
{
|
||||
name: "initialized notification",
|
||||
body: `{"jsonrpc":"2.0","method":"notifications/initialized"}`,
|
||||
wantStatus: http.StatusAccepted,
|
||||
},
|
||||
{
|
||||
name: "list tools",
|
||||
body: `{"jsonrpc":"2.0","id":2,"method":"tools/list"}`,
|
||||
wantStatus: http.StatusOK,
|
||||
wantBody: `"tools"`,
|
||||
},
|
||||
{
|
||||
name: "call a read-only tool",
|
||||
body: `{"jsonrpc":"2.0","id":3,"method":"tools/call","params":{"name":"fire_safety_search_place_candidates","arguments":{"place_name":"观水镇","limit":3}}}`,
|
||||
wantStatus: http.StatusOK,
|
||||
wantBody: `"structuredContent"`,
|
||||
},
|
||||
}
|
||||
|
||||
handler.ServeHTTP(response, request)
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
request := authenticatedMCPRequest(tt.body)
|
||||
request.Header.Set("MCP-Protocol-Version", "superagent-unconfigured-version")
|
||||
response := httptest.NewRecorder()
|
||||
|
||||
if response.Code != http.StatusOK || !strings.Contains(response.Body.String(), "MCP_PROTOCOL_VERSION_UNSUPPORTED") {
|
||||
t.Fatalf("status=%d body=%s", response.Code, response.Body.String())
|
||||
handler.ServeHTTP(response, request)
|
||||
|
||||
if response.Code != tt.wantStatus || (tt.wantBody != "" && !strings.Contains(response.Body.String(), tt.wantBody)) {
|
||||
t.Fatalf("status=%d body=%q", response.Code, response.Body.String())
|
||||
}
|
||||
})
|
||||
}
|
||||
if spatial.called != toolSearchPlaceCandidates {
|
||||
t.Fatalf("called=%q, want %q", spatial.called, toolSearchPlaceCandidates)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMCPInitializeAcceptsSuperAgentProtocolShapes(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
params string
|
||||
protocolHeader string
|
||||
}{
|
||||
{name: "missing params"},
|
||||
{name: "null params", params: "null"},
|
||||
{name: "missing version", params: `{}`},
|
||||
{name: "empty version", params: `{"protocolVersion":""}`},
|
||||
{name: "blank version", params: `{"protocolVersion":" "}`},
|
||||
{name: "non-string version", params: `{"protocolVersion":20250618}`},
|
||||
{name: "invalid params shape", params: `[]`},
|
||||
{name: "direct", params: `{"protocolVersion":"2025-06-18"}`},
|
||||
{name: "older", params: `{"protocolVersion":"2025-03-26"}`},
|
||||
{name: "newer and mismatched header", params: `{"protocolVersion":"2025-11-25"}`, protocolHeader: "2025-03-26"},
|
||||
{name: "unknown", params: `{"protocolVersion":"superagent-private-version"}`, protocolHeader: "superagent-header-version"},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
handler := newTestMCPHandler(t, &fakeMCPSpatialService{}, 4096, time.Second)
|
||||
body := `{"jsonrpc":"2.0","id":1,"method":"initialize"`
|
||||
if tt.params != "" {
|
||||
body += `,"params":` + tt.params
|
||||
}
|
||||
request := authenticatedMCPRequest(body + `}`)
|
||||
if tt.protocolHeader != "" {
|
||||
request.Header.Set("MCP-Protocol-Version", tt.protocolHeader)
|
||||
}
|
||||
response := httptest.NewRecorder()
|
||||
|
||||
handler.ServeHTTP(response, request)
|
||||
|
||||
var decoded rpcResponse
|
||||
if err := json.Unmarshal(response.Body.Bytes(), &decoded); err != nil {
|
||||
t.Fatalf("decode response: %v; body=%s", err, response.Body.String())
|
||||
}
|
||||
if response.Code != http.StatusOK ||
|
||||
!strings.Contains(response.Body.String(), `"protocolVersion":"2025-06-18"`) ||
|
||||
decoded.Error != nil {
|
||||
t.Fatalf("status=%d body=%s", response.Code, response.Body.String())
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestMCPProtocolHeaderDoesNotBlockSuperAgentRequests(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
header string
|
||||
}{
|
||||
{name: "missing"},
|
||||
{name: "server version", header: mcpProtocolVersion},
|
||||
{name: "older", header: "2025-03-26"},
|
||||
{name: "newer", header: "2025-11-25"},
|
||||
{name: "unknown", header: "superagent-private-version"},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
handler := newTestMCPHandler(t, &fakeMCPSpatialService{}, 4096, time.Second)
|
||||
request := authenticatedMCPRequest(`{"jsonrpc":"2.0","id":2,"method":"tools/list"}`)
|
||||
if tt.header == "" {
|
||||
request.Header.Del("MCP-Protocol-Version")
|
||||
} else {
|
||||
request.Header.Set("MCP-Protocol-Version", tt.header)
|
||||
}
|
||||
response := httptest.NewRecorder()
|
||||
|
||||
handler.ServeHTTP(response, request)
|
||||
|
||||
var decoded rpcResponse
|
||||
if err := json.Unmarshal(response.Body.Bytes(), &decoded); err != nil {
|
||||
t.Fatalf("decode response: %v; body=%s", err, response.Body.String())
|
||||
}
|
||||
if response.Code != http.StatusOK || decoded.Error != nil || !strings.Contains(response.Body.String(), `"tools"`) {
|
||||
t.Fatalf("status=%d body=%s", response.Code, response.Body.String())
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestMCPProtocolHeaderDoesNotBlockInitializedNotification(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
header string
|
||||
}{
|
||||
{name: "missing"},
|
||||
{name: "server version", header: mcpProtocolVersion},
|
||||
{name: "older", header: "2025-03-26"},
|
||||
{name: "newer", header: "2025-11-25"},
|
||||
{name: "unknown", header: "superagent-private-version"},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
handler := newTestMCPHandler(t, &fakeMCPSpatialService{}, 4096, time.Second)
|
||||
request := authenticatedMCPRequest(`{"jsonrpc":"2.0","method":"notifications/initialized"}`)
|
||||
if tt.header == "" {
|
||||
request.Header.Del("MCP-Protocol-Version")
|
||||
} else {
|
||||
request.Header.Set("MCP-Protocol-Version", tt.header)
|
||||
}
|
||||
response := httptest.NewRecorder()
|
||||
|
||||
handler.ServeHTTP(response, request)
|
||||
|
||||
if response.Code != http.StatusAccepted || response.Body.Len() != 0 {
|
||||
t.Fatalf("status=%d body=%q", response.Code, response.Body.String())
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestMCPProtocolHeaderDoesNotBlockToolCalls(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
header string
|
||||
}{
|
||||
{name: "missing"},
|
||||
{name: "server version", header: mcpProtocolVersion},
|
||||
{name: "older", header: "2025-03-26"},
|
||||
{name: "newer", header: "2025-11-25"},
|
||||
{name: "unknown", header: "superagent-private-version"},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
spatial := &fakeMCPSpatialService{}
|
||||
handler := newTestMCPHandler(t, spatial, 4096, time.Second)
|
||||
request := authenticatedMCPRequest(`{"jsonrpc":"2.0","id":3,"method":"tools/call","params":{"name":"fire_safety_resolve_incident_context","arguments":{"longitude":121.7,"latitude":37.2}}}`)
|
||||
if tt.header == "" {
|
||||
request.Header.Del("MCP-Protocol-Version")
|
||||
} else {
|
||||
request.Header.Set("MCP-Protocol-Version", tt.header)
|
||||
}
|
||||
response := httptest.NewRecorder()
|
||||
|
||||
handler.ServeHTTP(response, request)
|
||||
|
||||
var decoded rpcResponse
|
||||
if err := json.Unmarshal(response.Body.Bytes(), &decoded); err != nil {
|
||||
t.Fatalf("decode response: %v; body=%s", err, response.Body.String())
|
||||
}
|
||||
if response.Code != http.StatusOK || decoded.Error != nil || spatial.called != toolResolveIncidentContext {
|
||||
t.Fatalf("status=%d called=%q body=%s", response.Code, spatial.called, response.Body.String())
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestMCPInitializeLogsFiniteCompatibilityResultWithoutInput(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
params string
|
||||
wantResult string
|
||||
}{
|
||||
{name: "direct", params: `{"protocolVersion":"2025-06-18"}`, wantResult: "direct_success"},
|
||||
{name: "other", params: `{"protocolVersion":"untrusted-client-version"}`, wantResult: "compatibility_success"},
|
||||
{name: "missing", wantResult: "compatibility_success"},
|
||||
{name: "non-string", params: `{"protocolVersion":20250618,"extra":"sensitive-value"}`, wantResult: "compatibility_success"},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
var logs strings.Builder
|
||||
handler, err := NewMCPHandler(&fakeMCPSpatialService{}, MCPOptions{
|
||||
AuthToken: testMCPToken,
|
||||
MaxBodyBytes: 4096,
|
||||
ToolTimeout: time.Second,
|
||||
Logger: log.New(&logs, "", 0),
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("NewMCPHandler() error = %v", err)
|
||||
}
|
||||
body := `{"jsonrpc":"2.0","id":1,"method":"initialize"`
|
||||
if tt.params != "" {
|
||||
body += `,"params":` + tt.params
|
||||
}
|
||||
request := authenticatedMCPRequest(body + `}`)
|
||||
response := httptest.NewRecorder()
|
||||
|
||||
handler.ServeHTTP(response, request)
|
||||
|
||||
if !strings.Contains(logs.String(), "operation=initialize result="+tt.wantResult) {
|
||||
t.Fatalf("logs=%q, want result %s", logs.String(), tt.wantResult)
|
||||
}
|
||||
for _, sensitive := range []string{"untrusted-client-version", "sensitive-value", testMCPToken} {
|
||||
if strings.Contains(logs.String(), sensitive) {
|
||||
t.Fatalf("logs contain request input %q: %q", sensitive, logs.String())
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
Reference in new issue
Block a user