summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorPaul Buetow <paul@buetow.org>2026-06-19 20:50:09 +0300
committerPaul Buetow <paul@buetow.org>2026-06-19 20:50:09 +0300
commit152f32cf1b0a5dad8a0443130557f95ae9c4fb0a (patch)
tree27dba9164ceab76d8a1700c56d2729ad68adc2a5
parent3b0816f11504ac9ef8f55b91b99371fbb55bb040 (diff)
Fix MCP protocol test cleanup for 7k0
-rw-r--r--internal/hexaimcp/run_test.go121
1 files changed, 96 insertions, 25 deletions
diff --git a/internal/hexaimcp/run_test.go b/internal/hexaimcp/run_test.go
index f4e130b..4a25781 100644
--- a/internal/hexaimcp/run_test.go
+++ b/internal/hexaimcp/run_test.go
@@ -12,6 +12,7 @@ import (
"os"
"path/filepath"
"strings"
+ "sync"
"testing"
"time"
@@ -20,16 +21,29 @@ import (
"codeberg.org/snonux/hexai/internal/promptstore"
)
+const (
+ responseReadTimeout = 2 * time.Second
+ serverExitTimeout = 2 * time.Second
+ serverStillRunningTimeout = 100 * time.Millisecond
+)
+
// mockServerRunner implements ServerRunner for testing
type mockServerRunner struct {
runFunc func() error
}
+type serverExit struct {
+ done chan struct{}
+ err error
+}
+
type fullProtocolServer struct {
- stdin *io.PipeWriter
- stdout io.Reader
- done <-chan error
- cancel context.CancelFunc
+ stdin *io.PipeWriter
+ stdout io.Reader
+ exit *serverExit
+ cancel context.CancelFunc
+ closeStdin sync.Once
+ closeStdinErr error
}
func (m *mockServerRunner) Run(context.Context) error {
@@ -43,16 +57,15 @@ func (m *mockServerRunner) Run(context.Context) error {
func TestFullProtocolFlow(t *testing.T) {
tmpDir, promptsDir := setupFullProtocolConfig(t)
server := startFullProtocolServer(t, tmpDir, promptsDir)
- defer server.cancel()
writeJSONRPCLine(t, server.stdin, initializeRequest())
assertInitializeResponse(t, readJSONRPCLine(t, server.stdout))
- assertServerStillRunning(t, server.done)
+ assertServerStillRunning(t, server.exit)
- if err := server.stdin.Close(); err != nil {
+ if err := server.closeInput(); err != nil {
t.Fatalf("close stdin writer: %v", err)
}
- if err := waitForServer(t, server.done); err != nil {
+ if err := waitForServer(t, server.exit); err != nil {
t.Fatalf("server returned error: %v", err)
}
@@ -78,11 +91,11 @@ func setupFullProtocolConfig(t *testing.T) (string, string) {
return tmpDir, filepath.Join(tmpDir, "prompts")
}
-func startFullProtocolServer(t *testing.T, tmpDir, promptsDir string) fullProtocolServer {
+func startFullProtocolServer(t *testing.T, tmpDir, promptsDir string) *fullProtocolServer {
t.Helper()
stdinReader, stdinWriter := io.Pipe()
stdoutReader, stdoutWriter := io.Pipe()
- done := make(chan error, 1)
+ exit := newServerExit()
ctx, cancel := context.WithCancel(context.Background())
loadOpts := appconfig.LoadOptions{
@@ -96,20 +109,65 @@ func startFullProtocolServer(t *testing.T, tmpDir, promptsDir string) fullProtoc
}
go func() {
+ var runErr error
defer func() {
_ = stdinReader.Close()
_ = stdoutWriter.Close()
+ exit.finish(runErr)
}()
logPath := filepath.Join(tmpDir, "mcp.log")
- done <- runWithFactoryLoadOptions(ctx, logPath, loadOpts, overrides, stdinReader, stdoutWriter, &bytes.Buffer{}, serverFactory)
+ runErr = runWithFactoryLoadOptions(ctx, logPath, loadOpts, overrides, stdinReader, stdoutWriter, &bytes.Buffer{}, serverFactory)
}()
- return fullProtocolServer{
+ server := &fullProtocolServer{
stdin: stdinWriter,
stdout: stdoutReader,
- done: done,
+ exit: exit,
cancel: cancel,
}
+ t.Cleanup(func() {
+ _ = server.closeInput()
+ server.cancel()
+ defer func() {
+ _ = stdoutReader.Close()
+ }()
+
+ err, ok := server.exit.wait(serverExitTimeout)
+ if !ok {
+ t.Errorf("server did not exit during cleanup after %s", serverExitTimeout)
+ return
+ }
+ if err != nil {
+ t.Errorf("server cleanup returned error: %v", err)
+ }
+ })
+
+ return server
+}
+
+func newServerExit() *serverExit {
+ return &serverExit{done: make(chan struct{})}
+}
+
+func (e *serverExit) finish(err error) {
+ e.err = err
+ close(e.done)
+}
+
+func (e *serverExit) wait(timeout time.Duration) (error, bool) {
+ select {
+ case <-e.done:
+ return e.err, true
+ case <-time.After(timeout):
+ return nil, false
+ }
+}
+
+func (s *fullProtocolServer) closeInput() error {
+ s.closeStdin.Do(func() {
+ s.closeStdinErr = s.stdin.Close()
+ })
+ return s.closeStdinErr
}
func initializeRequest() map[string]any {
@@ -144,10 +202,27 @@ func writeJSONRPCLine(t *testing.T, w io.Writer, req map[string]any) {
func readJSONRPCLine(t *testing.T, r io.Reader) mcp.Response {
t.Helper()
- line, err := bufio.NewReader(r).ReadBytes('\n')
- if err != nil {
- t.Fatalf("read response: %v", err)
+ type readResult struct {
+ line []byte
+ err error
+ }
+ result := make(chan readResult, 1)
+ go func() {
+ line, err := bufio.NewReader(r).ReadBytes('\n')
+ result <- readResult{line: line, err: err}
+ }()
+
+ var line []byte
+ select {
+ case res := <-result:
+ if res.err != nil {
+ t.Fatalf("read response: %v", res.err)
+ }
+ line = res.line
+ case <-time.After(responseReadTimeout):
+ t.Fatalf("timed out reading response after %s", responseReadTimeout)
}
+
var resp mcp.Response
if err := json.Unmarshal(line, &resp); err != nil {
t.Fatalf("unmarshal response %q: %v", line, err)
@@ -168,23 +243,19 @@ func assertInitializeResponse(t *testing.T, resp mcp.Response) {
}
}
-func assertServerStillRunning(t *testing.T, done <-chan error) {
+func assertServerStillRunning(t *testing.T, exit *serverExit) {
t.Helper()
- select {
- case err := <-done:
+ if err, ok := exit.wait(serverStillRunningTimeout); ok {
t.Fatalf("server returned before stdin was closed: %v", err)
- default:
}
}
-func waitForServer(t *testing.T, done <-chan error) error {
+func waitForServer(t *testing.T, exit *serverExit) error {
t.Helper()
- select {
- case err := <-done:
+ if err, ok := exit.wait(serverExitTimeout); ok {
return err
- case <-time.After(2 * time.Second):
- t.Fatal("server did not exit after stdin closed")
}
+ t.Fatal("server did not exit after stdin closed")
return nil
}