snapshot of eee7b5cfc0776c68816d7fdaa8af9d7a0e0da2e6 Annotations about the code that implements koment.

internal/mcp/transport.go

1 package mcp
2
3 import (
4 "context"
5 "errors"
6 "flag"
7 "fmt"
8 "io"
9 "net"
10 "net/http"
11 "os"
12 "time"
13
14 sdk "github.com/modelcontextprotocol/go-sdk/mcp"
15
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 )
21
22 const (
23 shutdownGrace = 5 * time.Second
24 headerTimeout = 10 * time.Second
25 sweepInterval = 30 * time.Second
26 )
27
28 const transportUsage = `koment mcp serves annotations to agents.
29
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
34
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 `
39
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 }
48
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 }
68
69 repositories, err := loadRepositories()
70 if err != nil {
71 return err
72 }
73
74 ctx := context.Background()
75 recorder := startMetrics(ctx, repositories, *metricsAddress, stderr)
76
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 }
85
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 }
93
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 }
102
103 func startMetrics(ctx context.Context, repositories *repository.Set, address string, stderr io.Writer) metrics.Recorder {
104 if address == "" {
105 return metrics.Discard{}
106 }
107
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 }
130
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)
137
138 handler := sdk.NewStreamableHTTPHandler(
139 func(*http.Request) *sdk.Server { return newServer(repositories, recorder, false) },
140 &sdk.StreamableHTTPOptions{JSONResponse: jsonResponses},
141 )
142
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}
148
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 }()
157
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 }
164
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 }

Find an annotation

Search file paths, rationale, kinds, and authors.

moveEnter openEsc close