Files
ResolveSpec/pkg/resolvemcp/hardening_test.go
T

81 lines
2.5 KiB
Go

package resolvemcp
import (
"fmt"
"net/http"
"net/http/httptest"
"sync"
"testing"
"github.com/DATA-DOG/go-sqlmock"
)
func TestHookRegistryConcurrentUse(t *testing.T) {
r := NewHookRegistry()
var wg sync.WaitGroup
for i := 0; i < 8; i++ {
wg.Add(3)
go func() { defer wg.Done(); r.Register(BeforeRead, func(*HookContext) error { return nil }) }()
go func() { defer wg.Done(); _ = r.Execute(BeforeRead, &HookContext{}) }()
go func() { defer wg.Done(); _ = r.HasHooks(BeforeRead); r.Clear(AfterRead) }()
}
wg.Wait()
}
// Update and delete report a missing row with the same error, so ids cannot be probed.
func TestNotFoundErrorsAreUniform(t *testing.T) {
h, mock, ctx := newTxHarness(t)
empty := sqlmock.NewRows([]string{"id", "name"})
mock.ExpectBegin()
mock.ExpectQuery(`SELECT`).WillReturnRows(empty)
mock.ExpectRollback()
_, errU := h.executeUpdate(ctx, "public", "items", "9", map[string]interface{}{"name": "x"})
mock.ExpectBegin()
mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}))
mock.ExpectRollback()
_, errD := h.executeDelete(ctx, "public", "items", "9")
if errU == nil || errD == nil || errU.Error() != errD.Error() {
t.Fatalf("update %v / delete %v must be the same error", errU, errD)
}
}
func TestSSEHostAllowlistAndPoolCap(t *testing.T) {
h, _, _ := newTxHarness(t)
h.config.AllowedHosts = []string{"mcp.example.com"}
d := &dynamicSSEHandler{h: h}
r := httptest.NewRequest(http.MethodPost, "/mcp/message", nil)
r.Host = "evil.example.net"
w := httptest.NewRecorder()
d.ServeHTTP(w, r)
if w.Code != http.StatusBadRequest {
t.Errorf("foreign host: status %d, want 400", w.Code)
}
r = httptest.NewRequest(http.MethodPost, "/mcp/message", nil)
r.Host = "mcp.example.com"
r.Header.Set("X-Forwarded-Proto", "javascript")
w = httptest.NewRecorder()
d.ServeHTTP(w, r)
if w.Code != http.StatusBadRequest {
t.Errorf("bad proto: status %d, want 400", w.Code)
}
h.config.AllowedHosts = nil
for i := 0; i < maxSSEPool+5; i++ {
r = httptest.NewRequest(http.MethodPost, "/mcp/message?sessionId=x", nil)
r.Host = fmt.Sprintf("h%d.example.com", i)
d.ServeHTTP(httptest.NewRecorder(), r)
}
if len(d.pool) > maxSSEPool {
t.Errorf("pool grew to %d, cap is %d", len(d.pool), maxSSEPool)
}
r = httptest.NewRequest(http.MethodPost, "/mcp/message", nil)
r.Host = "one-more.example.com"
w = httptest.NewRecorder()
d.ServeHTTP(w, r)
if w.Code != http.StatusServiceUnavailable {
t.Errorf("full pool: status %d, want 503", w.Code)
}
}