internal/mcp/transport.go
1
package mcp
3
import (
4
"context"
5
"errors"
6
"flag"
7
"fmt"
8
"io"
9
"net"
10
"net/http"
11
"os"
12
"time"
14
sdk "github.com/modelcontextprotocol/go-sdk/mcp"
16
"github.com/janpuc/koment/internal/config"
17
"github.com/janpuc/koment/internal/listen"
18
"github.com/janpuc/koment/internal/metrics"
19
"github.com/janpuc/koment/internal/repository"
20
)
22
const (
23
shutdownGrace = 5 * time.Second
24
headerTimeout = 10 * time.Second
25
sweepInterval = 30 * time.Second
26
)
28
const transportUsage = `koment mcp serves annotations to agents.
30
koment mcp stdio (default)
31
koment mcp --write stdio with mutation tools
32
koment mcp --http <addr> HTTP, JSON responses
33
koment mcp --streamable-http <addr> HTTP, server-sent events
35
<addr> may be a bare port. A host is added if omitted, and it is the loopback
36
interface: the server has no authentication, so anything that can reach the port
37
can read every annotation in the repository.
38
`
40
// Serve parses the mcp subcommand's own flags, so that package cli never links
41
// the MCP SDK.
42
func Serve(args []string, stderr io.Writer) error {
43
flags := flag.NewFlagSet("mcp", flag.ContinueOnError)
44
flags.SetOutput(stderr)
45
flags.Usage = func() {
46
fmt.Fprint(stderr, transportUsage, "\nFlags (each also settable from the environment):\n", config.Usage(flags))
47
}
49
httpAddress := flags.String("http", "", "serve over HTTP at this address, with JSON responses")
50
streamableAddress := flags.String("streamable-http", "", "serve over HTTP at this address, with SSE responses")
51
metricsAddress := flags.String("metrics", "", "serve Prometheus metrics on this separate address; off unless given")
52
writes := flags.Bool("write", false, "register local mutation tools; valid only with stdio")
53
if err := flags.Parse(args); err != nil {
54
return err
55
}
56
if err := config.FromEnvironment(flags); err != nil {
57
return err
58
}
59
if flags.NArg() > 0 {
60
return fmt.Errorf("mcp takes no arguments, got %s", flags.Arg(0))
61
}
62
if *httpAddress != "" && *streamableAddress != "" {
63
return errors.New("--http and --streamable-http are alternatives; choose one")
64
}
65
if *writes && (*httpAddress != "" || *streamableAddress != "") {
66
return errors.New("--write is available only over stdio; unauthenticated HTTP never registers mutation tools")
67
}
69
repositories, err := loadRepositories()
70
if err != nil {
71
return err
72
}
74
ctx := context.Background()
75
recorder := startMetrics(ctx, repositories, *metricsAddress, stderr)
77
switch {
78
case *httpAddress != "":
79
return serveHTTP(ctx, repositories, *httpAddress, true, stderr, recorder)
80
case *streamableAddress != "":
81
return serveHTTP(ctx, repositories, *streamableAddress, false, stderr, recorder)
82
}
83
return newServer(repositories, recorder, *writes).Run(ctx, &sdk.StdioTransport{})
84
}
86
func loadRepositories() (*repository.Set, error) {
87
workingDirectory, err := os.Getwd()
88
if err != nil {
89
return nil, fmt.Errorf("finding the working directory: %w", err)
90
}
91
return repository.Load(workingDirectory)
92
}
94
func sweepAll(repositories *repository.Set, recorder metrics.Recorder) error {
95
for _, entry := range repositories.All() {
96
if err := metrics.Sweep(entry.Store(), recorder); err != nil {
97
return err
98
}
99
}
100
return nil
101
}
103
func startMetrics(ctx context.Context, repositories *repository.Set, address string, stderr io.Writer) metrics.Recorder {
104
if address == "" {
105
return metrics.Discard{}
106
}
108
recorder := metrics.New()
109
go func() {
110
if err := recorder.Serve(ctx, address, stderr); err != nil {
111
fmt.Fprintf(stderr, "koment: metrics: %v\n", err)
112
}
113
}()
114
go func() {
115
ticker := time.NewTicker(sweepInterval)
116
defer ticker.Stop()
117
for {
118
if err := sweepAll(repositories, recorder); err != nil {
119
fmt.Fprintf(stderr, "koment: metrics sweep: %v\n", err)
120
}
121
select {
122
case <-ctx.Done():
123
return
124
case <-ticker.C:
125
}
126
}
127
}()
128
return recorder
129
}
131
func serveHTTP(ctx context.Context, repositories *repository.Set, address string, jsonResponses bool, stderr io.Writer, recorder metrics.Recorder) error {
132
resolved, err := listen.Address(address)
133
if err != nil {
134
return err
135
}
136
listen.WarnIfPublic(resolved, stderr)
138
handler := sdk.NewStreamableHTTPHandler(
139
func(*http.Request) *sdk.Server { return newServer(repositories, recorder, false) },
140
&sdk.StreamableHTTPOptions{JSONResponse: jsonResponses},
141
)
143
listener, err := net.Listen("tcp", resolved)
144
if err != nil {
145
return fmt.Errorf("listening on %s: %w", resolved, err)
146
}
147
server := &http.Server{Handler: handler, ReadHeaderTimeout: headerTimeout}
149
go func() {
150
<-ctx.Done()
151
timeout, cancel := context.WithTimeout(context.Background(), shutdownGrace)
152
defer cancel()
153
if err := server.Shutdown(timeout); err != nil {
154
fmt.Fprintf(stderr, "koment: shutting down: %v\n", err)
155
}
156
}()
158
fmt.Fprintf(stderr, "koment: serving %d annotated files at http://%s\n", annotatedFileCount(repositories), listener.Addr())
159
if err := server.Serve(listener); !errors.Is(err, http.ErrServerClosed) {
160
return err
161
}
162
return nil
163
}
165
func annotatedFileCount(repositories *repository.Set) int {
166
total := 0
167
for _, entry := range repositories.All() {
168
files, err := entry.Store().AnnotatedFiles()
169
if err != nil {
170
continue
171
}
172
total += len(files)
173
}
174
return total
175
}