mirror of
https://github.com/bitechdev/ResolveSpec.git
synced 2026-10-02 19:41:57 +00:00
81 lines
2.5 KiB
Go
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)
|
|
}
|
|
}
|