diff --git a/docs/client.md b/docs/client.md index 23d57be1..d42c767c 100644 --- a/docs/client.md +++ b/docs/client.md @@ -546,3 +546,43 @@ that optional capabilities outside the core protocol can be declared on the wire. Keys are namespaced as `"{vendor-prefix}/{extension-name}"`; values are per-extension settings objects. +Use `AddExtension` to declare one and `HasExtension` to test for one: + +```go +caps := &mcp.ClientCapabilities{} +caps.AddExtension("io.example/my-extension", nil) +client := mcp.NewClient(impl, &mcp.ClientOptions{Capabilities: caps}) + +cs, err := client.Connect(ctx, transport, nil) +... +if cs.InitializeResult().Capabilities.HasExtension("io.example/my-extension") { + // The server declared it too. +} +``` + +#### Tasks + +[SEP-2663](https://github.com/modelcontextprotocol/modelcontextprotocol/blob/main/seps/2663-tasks-extension.md) +moved tasks out of the core protocol and into the +[tasks extension](https://github.com/modelcontextprotocol/ext-tasks/blob/main/specification/draft/tasks.md), +identified by the `mcp.ExtensionTasks` constant. A server that has negotiated +it may answer a request with a durable task handle instead of the result that +was asked for, which the client then polls to completion. + +**The SDK does not implement task execution**, and does not declare the +extension by default. Declaring it is a promise to the peer: a client that +declares it must be prepared for any eligible request to return a task handle +instead of a result. Only declare it if you implement that polling flow +yourself. + +If a server returns a task handle anyway, decoding fails with +`*mcp.UnsupportedTaskResultError`, which carries the task ID: + +```go +res, err := cs.CallTool(ctx, params) +var terr *mcp.UnsupportedTaskResultError +if errors.As(err, &terr) { + log.Printf("server created task %s, which this SDK cannot resolve", terr.TaskID) +} +``` + diff --git a/docs/server.md b/docs/server.md index 8973e3ae..d777f0ab 100644 --- a/docs/server.md +++ b/docs/server.md @@ -1203,6 +1203,36 @@ capabilities outside the core protocol can be declared on the wire. Keys are namespaced as `"{vendor-prefix}/{extension-name}"`; values are per-extension settings objects. +Use `AddExtension` to declare one, and `HasExtension` to test what the client +declared. Client capabilities are read from the request, since as of protocol +version 2026-07-28 they travel in each request's `_meta` rather than in the +initialize handshake: + +```go +caps := &mcp.ServerCapabilities{} +caps.AddExtension("io.example/my-extension", nil) +server := mcp.NewServer(impl, &mcp.ServerOptions{Capabilities: caps}) + +// Inside a tool handler: +if req.ClientCapabilities().HasExtension("io.example/my-extension") { + // The client declared it too. +} +``` + +#### Tasks + +[SEP-2663](https://github.com/modelcontextprotocol/modelcontextprotocol/blob/main/seps/2663-tasks-extension.md) +moved tasks out of the core protocol and into the +[tasks extension](https://github.com/modelcontextprotocol/ext-tasks/blob/main/specification/draft/tasks.md), +identified by the `mcp.ExtensionTasks` constant. A server that has negotiated +it may answer a request with a durable task handle instead of the result that +was asked for, which the client then polls to completion. + +**The SDK does not implement task execution**, and does not declare the +extension by default. Declaring it is a promise to the peer: a server that +declares it must serve `tasks/get`, `tasks/update` and `tasks/cancel`. Only +declare it if you implement those yourself. + ### Pagination Server-side feature lists may be diff --git a/internal/docs/client.src.md b/internal/docs/client.src.md index c0f64b4b..be806de0 100644 --- a/internal/docs/client.src.md +++ b/internal/docs/client.src.md @@ -235,3 +235,43 @@ that optional capabilities outside the core protocol can be declared on the wire. Keys are namespaced as `"{vendor-prefix}/{extension-name}"`; values are per-extension settings objects. +Use `AddExtension` to declare one and `HasExtension` to test for one: + +```go +caps := &mcp.ClientCapabilities{} +caps.AddExtension("io.example/my-extension", nil) +client := mcp.NewClient(impl, &mcp.ClientOptions{Capabilities: caps}) + +cs, err := client.Connect(ctx, transport, nil) +... +if cs.InitializeResult().Capabilities.HasExtension("io.example/my-extension") { + // The server declared it too. +} +``` + +#### Tasks + +[SEP-2663](https://github.com/modelcontextprotocol/modelcontextprotocol/blob/main/seps/2663-tasks-extension.md) +moved tasks out of the core protocol and into the +[tasks extension](https://github.com/modelcontextprotocol/ext-tasks/blob/main/specification/draft/tasks.md), +identified by the `mcp.ExtensionTasks` constant. A server that has negotiated +it may answer a request with a durable task handle instead of the result that +was asked for, which the client then polls to completion. + +**The SDK does not implement task execution**, and does not declare the +extension by default. Declaring it is a promise to the peer: a client that +declares it must be prepared for any eligible request to return a task handle +instead of a result. Only declare it if you implement that polling flow +yourself. + +If a server returns a task handle anyway, decoding fails with +`*mcp.UnsupportedTaskResultError`, which carries the task ID: + +```go +res, err := cs.CallTool(ctx, params) +var terr *mcp.UnsupportedTaskResultError +if errors.As(err, &terr) { + log.Printf("server created task %s, which this SDK cannot resolve", terr.TaskID) +} +``` + diff --git a/internal/docs/server.src.md b/internal/docs/server.src.md index ff663836..6c14f9d0 100644 --- a/internal/docs/server.src.md +++ b/internal/docs/server.src.md @@ -517,6 +517,36 @@ capabilities outside the core protocol can be declared on the wire. Keys are namespaced as `"{vendor-prefix}/{extension-name}"`; values are per-extension settings objects. +Use `AddExtension` to declare one, and `HasExtension` to test what the client +declared. Client capabilities are read from the request, since as of protocol +version 2026-07-28 they travel in each request's `_meta` rather than in the +initialize handshake: + +```go +caps := &mcp.ServerCapabilities{} +caps.AddExtension("io.example/my-extension", nil) +server := mcp.NewServer(impl, &mcp.ServerOptions{Capabilities: caps}) + +// Inside a tool handler: +if req.ClientCapabilities().HasExtension("io.example/my-extension") { + // The client declared it too. +} +``` + +#### Tasks + +[SEP-2663](https://github.com/modelcontextprotocol/modelcontextprotocol/blob/main/seps/2663-tasks-extension.md) +moved tasks out of the core protocol and into the +[tasks extension](https://github.com/modelcontextprotocol/ext-tasks/blob/main/specification/draft/tasks.md), +identified by the `mcp.ExtensionTasks` constant. A server that has negotiated +it may answer a request with a durable task handle instead of the result that +was asked for, which the client then polls to completion. + +**The SDK does not implement task execution**, and does not declare the +extension by default. Declaring it is a promise to the peer: a server that +declares it must serve `tasks/get`, `tasks/update` and `tasks/cancel`. Only +declare it if you implement those yourself. + ### Pagination Server-side feature lists may be diff --git a/mcp/protocol.go b/mcp/protocol.go index d5ef5400..2a26a674 100644 --- a/mcp/protocol.go +++ b/mcp/protocol.go @@ -26,6 +26,11 @@ const ( // input before it can complete the request. The client should fulfill the // InputRequests and retry the call with the responses. resultTypeInputRequired resultType = "input_required" + + // resultTypeTask is reserved by the io.modelcontextprotocol/tasks + // extension to discriminate a CreateTaskResult from a standard result. + // See [ExtensionTasks]. + resultTypeTask resultType = "task" ) type completeResultWithType struct { @@ -470,10 +475,14 @@ func (x *CallToolResult) UnmarshalJSON(data []byte) error { Content []*wireContent `json:"content"` StructuredContent json.RawMessage `json:"structuredContent"` ResultType resultType `json:"resultType"` + TaskID string `json:"taskId"` } if err := internaljson.Unmarshal(data, &wire); err != nil { return err } + if wire.ResultType == resultTypeTask { + return &UnsupportedTaskResultError{TaskID: wire.TaskID} + } if len(wire.StructuredContent) > 0 { unmarshal := internaljson.UnmarshalUseNumber if structuredcontentfloat64 == "1" { @@ -598,6 +607,16 @@ func (c *ClientCapabilities) AddExtension(name string, settings map[string]any) c.Extensions[name] = settings } +// HasExtension reports whether c declares the extension with the given name. +// It is safe to call on a nil *ClientCapabilities. +func (c *ClientCapabilities) HasExtension(name string) bool { + if c == nil { + return false + } + _, ok := c.Extensions[name] + return ok +} + // clone returns a copy of the ClientCapabilities. // Values in the Extensions and Experimental maps are shallow-copied. func (c *ClientCapabilities) clone() *ClientCapabilities { @@ -1116,10 +1135,14 @@ func (x *GetPromptResult) UnmarshalJSON(data []byte) error { var wire struct { res ResultType resultType `json:"resultType"` + TaskID string `json:"taskId"` } if err := internaljson.Unmarshal(data, &wire); err != nil { return err } + if wire.ResultType == resultTypeTask { + return &UnsupportedTaskResultError{TaskID: wire.TaskID} + } wire.res.resultType = wire.ResultType *x = GetPromptResult(wire.res) return nil @@ -1747,10 +1770,14 @@ func (x *ReadResourceResult) UnmarshalJSON(data []byte) error { var wire struct { res ResultType resultType `json:"resultType"` + TaskID string `json:"taskId"` } if err := internaljson.Unmarshal(data, &wire); err != nil { return err } + if wire.ResultType == resultTypeTask { + return &UnsupportedTaskResultError{TaskID: wire.TaskID} + } wire.res.resultType = wire.ResultType *x = ReadResourceResult(wire.res) return nil @@ -2406,6 +2433,16 @@ func (c *ServerCapabilities) AddExtension(name string, settings map[string]any) c.Extensions[name] = settings } +// HasExtension reports whether c declares the extension with the given name. +// It is safe to call on a nil *ServerCapabilities. +func (c *ServerCapabilities) HasExtension(name string) bool { + if c == nil { + return false + } + _, ok := c.Extensions[name] + return ok +} + // clone returns a copy of the ServerCapabilities. // Values in the Extensions and Experimental maps are shallow-copied. func (c *ServerCapabilities) clone() *ServerCapabilities { diff --git a/mcp/tasks.go b/mcp/tasks.go new file mode 100644 index 00000000..16380cc5 --- /dev/null +++ b/mcp/tasks.go @@ -0,0 +1,88 @@ +// Copyright 2025 The Go MCP SDK Authors. All rights reserved. +// Use of this source code is governed by the license +// that can be found in the LICENSE file. + +package mcp + +import ( + "crypto/rand" + "fmt" +) + +// ExtensionTasks identifies the MCP Tasks extension, which lets a server answer +// a request with a durable task handle instead of the request's normal result. +// +// This SDK does not implement task execution, and does not declare the +// extension by default: declaring it obliges a client to poll a task handle to +// completion, and a server to serve the tasks/* methods. +// +// See https://github.com/modelcontextprotocol/ext-tasks/blob/main/specification/draft/tasks.md. +const ExtensionTasks = "io.modelcontextprotocol/tasks" + +// UnsupportedTaskResultError reports that a peer answered a request with a task +// handle from the [ExtensionTasks] extension, which this SDK cannot resolve. +type UnsupportedTaskResultError struct { + // TaskID identifies the created task, for manual polling or cancellation. + TaskID string +} + +func (e *UnsupportedTaskResultError) Error() string { + return fmt.Sprintf("peer created task %q: the %s extension is not implemented", e.TaskID, ExtensionTasks) +} + +// TaskStatus is the state of a task in the [ExtensionTasks] extension. +// Values outside the constants below are preserved: the extension may add +// statuses, and a peer's status must round-trip unchanged. +// +// See https://github.com/modelcontextprotocol/ext-tasks/blob/main/specification/draft/tasks.md. +type TaskStatus string + +const ( + // TaskStatusWorking means the request is currently being processed. + TaskStatusWorking TaskStatus = "working" + // TaskStatusInputRequired means the server needs input from the client. + TaskStatusInputRequired TaskStatus = "input_required" + // TaskStatusCompleted means the request finished and its result is available. + TaskStatusCompleted TaskStatus = "completed" + // TaskStatusCancelled means the request was cancelled before completion. + TaskStatusCancelled TaskStatus = "cancelled" + // TaskStatusFailed means the request failed with a JSON-RPC error. + TaskStatusFailed TaskStatus = "failed" +) + +// Task is the operational metadata for a task in the [ExtensionTasks] extension. +// Derived shapes that carry inputRequests, result, or error are not modeled +// here; those belong to task execution, which this SDK does not implement. +// +// Timestamps are strings, matching [Annotations.LastModified]: the extension +// types them as ISO 8601 text, and parsing them as time.Time would reject +// forms the spec allows and rewrite the value on the way back out. +// +// See https://github.com/modelcontextprotocol/ext-tasks/blob/main/specification/draft/tasks.md. +type Task struct { + // TaskID is the server-generated identifier for this task. + TaskID string `json:"taskId"` + // Status is the current task state. + Status TaskStatus `json:"status"` + // StatusMessage is an optional description of the current state. + StatusMessage string `json:"statusMessage,omitempty"` + // CreatedAt is the ISO 8601 timestamp when the task was created. + CreatedAt string `json:"createdAt"` + // LastUpdatedAt is the ISO 8601 timestamp when the task was last updated. + LastUpdatedAt string `json:"lastUpdatedAt"` + // TTLMs is the time-to-live from creation, in milliseconds. + // Nil encodes as JSON null, which the extension defines as unlimited. + // The field is required, so a nil pointer is sent as null. int64 holds + // a TTL past the 32-bit range: a year is about 3.15e10 ms. + TTLMs *int64 `json:"ttlMs"` + // PollIntervalMs is the suggested polling interval in milliseconds. + // Omitted when unset. + PollIntervalMs *int64 `json:"pollIntervalMs,omitempty"` +} + +// newTaskID returns an unguessable task ID. The extension requires +// server-generated IDs with enough entropy that a third party cannot +// enumerate them. [crypto/rand.Text] supplies at least 128 bits. +func newTaskID() string { + return rand.Text() +} diff --git a/mcp/tasks_test.go b/mcp/tasks_test.go new file mode 100644 index 00000000..d284c9a0 --- /dev/null +++ b/mcp/tasks_test.go @@ -0,0 +1,387 @@ +// Copyright 2025 The Go MCP SDK Authors. All rights reserved. +// Use of this source code is governed by the license +// that can be found in the LICENSE file. + +package mcp + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "sync/atomic" + "testing" + + "github.com/google/go-cmp/cmp" + "github.com/google/jsonschema-go/jsonschema" +) + +func TestHasExtension(t *testing.T) { + tests := []struct { + name string + extensions map[string]any + lookup string + want bool + }{ + {"nil map", nil, ExtensionTasks, false}, + {"empty map", map[string]any{}, ExtensionTasks, false}, + {"absent", map[string]any{"io.example/other": map[string]any{}}, ExtensionTasks, false}, + {"present, empty settings", map[string]any{ExtensionTasks: map[string]any{}}, ExtensionTasks, true}, + {"present, with settings", map[string]any{ExtensionTasks: map[string]any{"k": "v"}}, ExtensionTasks, true}, + {"present, nil settings", map[string]any{ExtensionTasks: nil}, ExtensionTasks, true}, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + client := &ClientCapabilities{Extensions: tc.extensions} + if got := client.HasExtension(tc.lookup); got != tc.want { + t.Errorf("ClientCapabilities.HasExtension(%q) = %v, want %v", tc.lookup, got, tc.want) + } + server := &ServerCapabilities{Extensions: tc.extensions} + if got := server.HasExtension(tc.lookup); got != tc.want { + t.Errorf("ServerCapabilities.HasExtension(%q) = %v, want %v", tc.lookup, got, tc.want) + } + }) + } + + t.Run("nil receiver", func(t *testing.T) { + var client *ClientCapabilities + if client.HasExtension(ExtensionTasks) { + t.Error("(*ClientCapabilities)(nil).HasExtension = true, want false") + } + var server *ServerCapabilities + if server.HasExtension(ExtensionTasks) { + t.Error("(*ServerCapabilities)(nil).HasExtension = true, want false") + } + }) + + t.Run("round trip with AddExtension", func(t *testing.T) { + client := new(ClientCapabilities) + client.AddExtension(ExtensionTasks, nil) + if !client.HasExtension(ExtensionTasks) { + t.Error("ClientCapabilities.HasExtension after AddExtension = false, want true") + } + server := new(ServerCapabilities) + server.AddExtension(ExtensionTasks, nil) + if !server.HasExtension(ExtensionTasks) { + t.Error("ServerCapabilities.HasExtension after AddExtension = false, want true") + } + }) +} + +// TestServerSeesClientTasksExtension checks that a client declaring the tasks +// extension is visible to a server request handler, across both capability +// transports: the per-request _meta of protocol 2026-07-28, and the +// initialize handshake of older versions. +// +// The declared=false cases also pin that the SDK never declares the extension +// on the user's behalf, since it does not implement task execution. +func TestServerSeesClientTasksExtension(t *testing.T) { + for _, version := range []string{protocolVersion20260728, protocolVersion20251125} { + t.Run(version, func(t *testing.T) { + for _, declare := range []bool{true, false} { + t.Run(fmt.Sprintf("declared=%t", declare), func(t *testing.T) { + ctx := context.Background() + + var got atomic.Bool + server := NewServer(testImpl, nil) + server.AddTool( + &Tool{Name: "probe", InputSchema: &jsonschema.Schema{Type: "object"}}, + func(ctx context.Context, req *CallToolRequest) (*CallToolResult, error) { + got.Store(req.ClientCapabilities().HasExtension(ExtensionTasks)) + return &CallToolResult{Content: []Content{&TextContent{Text: "ok"}}}, nil + }) + + clientOpts := new(ClientOptions) + if declare { + caps := new(ClientCapabilities) + caps.AddExtension(ExtensionTasks, nil) + clientOpts.Capabilities = caps + } + + ct, st := NewInMemoryTransports() + ss, err := server.Connect(ctx, st, nil) + if err != nil { + t.Fatalf("server Connect: %v", err) + } + defer ss.Close() + + cs, err := NewClient(testImpl, clientOpts).Connect(ctx, ct, &ClientSessionOptions{ProtocolVersion: version}) + if err != nil { + t.Fatalf("client Connect: %v", err) + } + defer cs.Close() + + if _, err := cs.CallTool(ctx, &CallToolParams{Name: "probe"}); err != nil { + t.Fatalf("CallTool: %v", err) + } + if got.Load() != declare { + t.Errorf("handler saw %s = %v, want %v", ExtensionTasks, got.Load(), declare) + } + }) + } + }) + } +} + +// TestClientSeesServerTasksExtension checks that a server declaring the tasks +// extension is visible to the client, both through server/discover and through +// the legacy initialize handshake. +func TestClientSeesServerTasksExtension(t *testing.T) { + for _, version := range []string{protocolVersion20260728, protocolVersion20251125} { + t.Run(version, func(t *testing.T) { + for _, declare := range []bool{true, false} { + t.Run(fmt.Sprintf("declared=%t", declare), func(t *testing.T) { + ctx := context.Background() + + serverOpts := new(ServerOptions) + if declare { + caps := new(ServerCapabilities) + caps.AddExtension(ExtensionTasks, nil) + serverOpts.Capabilities = caps + } + + ct, st := NewInMemoryTransports() + ss, err := NewServer(testImpl, serverOpts).Connect(ctx, st, nil) + if err != nil { + t.Fatalf("server Connect: %v", err) + } + defer ss.Close() + + cs, err := NewClient(testImpl, nil).Connect(ctx, ct, &ClientSessionOptions{ProtocolVersion: version}) + if err != nil { + t.Fatalf("client Connect: %v", err) + } + defer cs.Close() + + if got := cs.InitializeResult().Capabilities.HasExtension(ExtensionTasks); got != declare { + t.Errorf("client saw %s = %v, want %v", ExtensionTasks, got, declare) + } + }) + } + }) + } +} + +// TestUnmarshalTaskResult checks that a CreateTaskResult from the tasks +// extension is rejected rather than silently decoding into an empty result. +func TestUnmarshalTaskResult(t *testing.T) { + const taskResult = `{ + "resultType": "task", + "taskId": "786512e2", + "status": "working", + "createdAt": "2026-01-01T00:00:00Z", + "lastUpdatedAt": "2026-01-01T00:00:00Z", + "ttlMs": 60000, + "pollIntervalMs": 5000 + }` + + targets := []struct { + name string + newTarget func() any + }{ + {"CallToolResult", func() any { return new(CallToolResult) }}, + {"GetPromptResult", func() any { return new(GetPromptResult) }}, + {"ReadResourceResult", func() any { return new(ReadResourceResult) }}, + } + + for _, target := range targets { + t.Run(target.name, func(t *testing.T) { + t.Run("task is rejected", func(t *testing.T) { + err := json.Unmarshal([]byte(taskResult), target.newTarget()) + var terr *UnsupportedTaskResultError + if !errors.As(err, &terr) { + t.Fatalf("Unmarshal error = %v, want *UnsupportedTaskResultError", err) + } + if got, want := terr.TaskID, "786512e2"; got != want { + t.Errorf("TaskID = %q, want %q", got, want) + } + }) + + t.Run("complete still decodes", func(t *testing.T) { + if err := json.Unmarshal([]byte(`{"resultType":"complete"}`), target.newTarget()); err != nil { + t.Errorf("Unmarshal of a complete result failed: %v", err) + } + }) + }) + } +} + +// taskResultStub is a [Result] that marshals to a CreateTaskResult, standing in +// for a server that has decided to answer a request with a task handle. +type taskResultStub struct { + ResultBase + taskID string +} + +func (s *taskResultStub) MarshalJSON() ([]byte, error) { + return json.Marshal(map[string]any{ + "resultType": "task", + "taskId": s.taskID, + "status": "working", + "createdAt": "2026-01-01T00:00:00Z", + "lastUpdatedAt": "2026-01-01T00:00:00Z", + "ttlMs": 60000, + "pollIntervalMs": 5000, + }) +} + +// TestCallToolTaskResultEndToEnd checks that the decode guard survives the real +// client call path with its error identity intact, rather than surfacing as an +// empty but successful tool result. +func TestCallToolTaskResultEndToEnd(t *testing.T) { + ctx := context.Background() + + server := NewServer(testImpl, nil) + server.AddTool( + &Tool{Name: "probe", InputSchema: &jsonschema.Schema{Type: "object"}}, + func(ctx context.Context, req *CallToolRequest) (*CallToolResult, error) { + return nil, errors.New("unreachable: intercepted by middleware") + }) + server.AddReceivingMiddleware(func(next MethodHandler) MethodHandler { + return func(ctx context.Context, method string, req Request) (Result, error) { + if method == methodCallTool { + return &taskResultStub{taskID: "786512e2"}, nil + } + return next(ctx, method, req) + } + }) + + ct, st := NewInMemoryTransports() + ss, err := server.Connect(ctx, st, nil) + if err != nil { + t.Fatalf("server Connect: %v", err) + } + defer ss.Close() + + cs, err := NewClient(testImpl, nil).Connect(ctx, ct, nil) + if err != nil { + t.Fatalf("client Connect: %v", err) + } + defer cs.Close() + + res, err := cs.CallTool(ctx, &CallToolParams{Name: "probe"}) + if err == nil { + t.Fatalf("CallTool succeeded with %+v, want an error", res) + } + var terr *UnsupportedTaskResultError + if !errors.As(err, &terr) { + t.Fatalf("CallTool error = %v, want *UnsupportedTaskResultError", err) + } + if got, want := terr.TaskID, "786512e2"; got != want { + t.Errorf("TaskID = %q, want %q", got, want) + } +} + +func TestTaskJSON(t *testing.T) { + ttl := int64(60000) + poll := int64(5000) + + tests := []struct { + name string + in Task + want string + }{ + { + name: "every field", + in: Task{ + TaskID: "786512e2-9e0d-44bd-8f29-789f320fe840", + Status: TaskStatusWorking, + StatusMessage: "The operation is now in progress.", + CreatedAt: "2025-11-25T10:30:00Z", + LastUpdatedAt: "2025-11-25T10:40:00Z", + TTLMs: &ttl, + PollIntervalMs: &poll, + }, + want: `{"taskId":"786512e2-9e0d-44bd-8f29-789f320fe840","status":"working","statusMessage":"The operation is now in progress.","createdAt":"2025-11-25T10:30:00Z","lastUpdatedAt":"2025-11-25T10:40:00Z","ttlMs":60000,"pollIntervalMs":5000}`, + }, + { + name: "ttlMs null", + in: Task{ + TaskID: "id", + Status: TaskStatusCompleted, + CreatedAt: "2025-11-25T10:30:00Z", + LastUpdatedAt: "2025-11-25T10:40:00Z", + }, + want: `{"taskId":"id","status":"completed","createdAt":"2025-11-25T10:30:00Z","lastUpdatedAt":"2025-11-25T10:40:00Z","ttlMs":null}`, + }, + { + name: "pollIntervalMs omitted", + in: Task{ + TaskID: "id", + Status: TaskStatusFailed, + CreatedAt: "2025-11-25T10:30:00Z", + LastUpdatedAt: "2025-11-25T10:40:00Z", + TTLMs: &ttl, + }, + want: `{"taskId":"id","status":"failed","createdAt":"2025-11-25T10:30:00Z","lastUpdatedAt":"2025-11-25T10:40:00Z","ttlMs":60000}`, + }, + { + name: "unknown status preserved", + in: Task{ + TaskID: "id", + Status: "queued", + CreatedAt: "2025-11-25T10:30:00Z", + LastUpdatedAt: "2025-11-25T10:40:00Z", + }, + want: `{"taskId":"id","status":"queued","createdAt":"2025-11-25T10:30:00Z","lastUpdatedAt":"2025-11-25T10:40:00Z","ttlMs":null}`, + }, + } + + for _, status := range []TaskStatus{ + TaskStatusWorking, + TaskStatusInputRequired, + TaskStatusCompleted, + TaskStatusCancelled, + TaskStatusFailed, + } { + tests = append(tests, struct { + name string + in Task + want string + }{ + name: "status " + string(status), + in: Task{ + TaskID: "id", + Status: status, + CreatedAt: "2025-11-25T10:30:00Z", + LastUpdatedAt: "2025-11-25T10:40:00Z", + TTLMs: &ttl, + }, + want: `{"taskId":"id","status":"` + string(status) + `","createdAt":"2025-11-25T10:30:00Z","lastUpdatedAt":"2025-11-25T10:40:00Z","ttlMs":60000}`, + }) + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + got, err := json.Marshal(tc.in) + if err != nil { + t.Fatalf("Marshal: %v", err) + } + if diff := cmp.Diff(tc.want, string(got)); diff != "" { + t.Errorf("Marshal mismatch (-want +got):\n%s", diff) + } + var out Task + if err := json.Unmarshal(got, &out); err != nil { + t.Fatalf("Unmarshal: %v", err) + } + if diff := cmp.Diff(tc.in, out); diff != "" { + t.Errorf("round trip mismatch (-want +got):\n%s", diff) + } + }) + } +} + +func TestNewTaskID(t *testing.T) { + seen := make(map[string]bool) + for i := 0; i < 32; i++ { + id := newTaskID() + if id == "" { + t.Fatal("newTaskID returned an empty ID") + } + if seen[id] { + t.Fatalf("newTaskID returned duplicate %q", id) + } + seen[id] = true + } +}