diff --git a/.github/workflows/dotnet-sdk-tests.yml b/.github/workflows/dotnet-sdk-tests.yml index f12f53bd96..2e35a645f7 100644 --- a/.github/workflows/dotnet-sdk-tests.yml +++ b/.github/workflows/dotnet-sdk-tests.yml @@ -42,6 +42,14 @@ jobs: - name: Build SDK run: dotnet build --no-restore + - name: Test structured output with reflection enabled + env: + DOTNET_ROLL_FORWARD: Major + run: >- + dotnet test test/GitHub.Copilot.SDK.Test.csproj --no-restore --framework net8.0 + -p:JsonSerializerIsReflectionEnabledByDefault=true + --filter "FullyQualifiedName~GitHub.Copilot.Test.Unit.ClientSessionLifetimeTests.StructuredOutput" + test: name: ".NET SDK Tests (${{ matrix.os }}, ${{ matrix.transport }}, ${{ matrix.backend }}, ${{ matrix.shard }})" if: github.event.repository.fork == false diff --git a/.github/workflows/rust-sdk-tests.yml b/.github/workflows/rust-sdk-tests.yml index 440641bbdc..b456dc0501 100644 --- a/.github/workflows/rust-sdk-tests.yml +++ b/.github/workflows/rust-sdk-tests.yml @@ -80,7 +80,7 @@ jobs: # embed it. Tests exec against the setup-copilot CLI via # COPILOT_CLI_PATH (the env override wins over the dev cache). # The dedicated `bundle` job below exercises the embed pipeline. - run: cargo test --no-default-features --features test-support -- --test-threads=4 --nocapture + run: cargo test --no-default-features --features test-support,derive -- --test-threads=4 --nocapture clippy: name: "Rust SDK Format and Clippy" @@ -139,7 +139,7 @@ jobs: - name: cargo clippy env: BUNDLED_CLI_CACHE_DIR: ${{ github.workspace }}/rust/.bundled-cli-cache - run: cargo clippy --all-targets --features test-support,bundled-in-process -- --no-deps -D warnings -D clippy::unwrap_used -D clippy::disallowed_macros -D clippy::await_holding_invalid_type + run: cargo clippy --all-targets --features test-support,bundled-in-process,derive -- --no-deps -D warnings -D clippy::unwrap_used -D clippy::disallowed_macros -D clippy::await_holding_invalid_type doc: name: "Rust SDK Docs" @@ -265,7 +265,7 @@ jobs: # The harness forces serial execution in-process (both the async semaphore and # libtest via --test-threads=1) because it mirrors each test's environment onto # the shared process environment, so RUST_E2E_CONCURRENCY is not set here. - run: cargo test --no-default-features --features test-support,bundled-in-process --test e2e -- --test-threads=1 --nocapture + run: cargo test --no-default-features --features test-support,bundled-in-process,derive --test e2e -- --test-threads=1 --nocapture # Validates the bundled-CLI build path on all three supported # platforms. While the regular `cargo test` job above also exercises @@ -376,8 +376,8 @@ jobs: export CARGO_TARGET_DIR=/tmp/copilot-sdk-rust-target if [ "$COPILOT_SDK_TEST_TRANSPORT" = "inprocess" ]; then unset RUST_E2E_CONCURRENCY - cargo test --no-default-features --features test-support,bundled-in-process --test e2e -- --test-threads=1 --nocapture + cargo test --no-default-features --features test-support,bundled-in-process,derive --test e2e -- --test-threads=1 --nocapture else export RUST_E2E_CONCURRENCY=4 - cargo test --no-default-features --features test-support -- --test-threads=4 --nocapture + cargo test --no-default-features --features test-support,derive -- --test-threads=4 --nocapture fi diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md index 4e7a3ee1e9..635d52be7e 100644 --- a/CONTRIBUTING.md +++ b/CONTRIBUTING.md @@ -58,6 +58,98 @@ Setup, build, and test instructions are maintained with each SDK: - [Rust](rust/README.md#development) - [Java](java/README.md#development-setup) +### Testing an unreleased runtime API + +The runtime's Rust contracts under `src/native/sdk-contract` produce both +`generated/api.schema.json` (RPC methods) and +`generated/session-events.schema.json` (event payloads). In a local checkout of +`github/copilot-agent-runtime`, build the runtime and emit these schemas: + +```bash +pnpm run build +pnpm bazel build //src/native/schema-codegen:schema-codegen +bazel-bin/src/native/schema-codegen/schema-codegen emit \ + --api "$PWD/generated/api.schema.json" \ + --session-events "$PWD/generated/session-events.schema.json" +``` + +The SDK generators normally download schemas from the pinned CLI release. To +use the local schemas instead, pass the event-schema path followed by the +RPC-schema path. From this repository's `scripts/codegen` directory: + +```bash +npm ci +for language in typescript csharp python go rust; do + node --import tsx "$language.ts" \ + "$RUNTIME_ROOT/generated/session-events.schema.json" \ + "$RUNTIME_ROOT/generated/api.schema.json" +done +``` + +Set `RUNTIME_ROOT` to the absolute path of the runtime checkout. Java's generator +at `java/scripts/codegen/java.ts` reads these files from +`java/scripts/codegen/target/schemas` instead of accepting positional arguments; +stage the local schemas there before running it. Do not hand-edit generated +wrappers. Regenerating against a newer runtime +also includes any other contract changes since the SDK's pinned release. + +Set `COPILOT_CLI_PATH` to the built runtime's `dist-cli/index.js` to run SDK E2Es +against that checkout rather than the packaged runtime. For example: + +```bash +export COPILOT_CLI_PATH="$RUNTIME_ROOT/dist-cli/index.js" +# Supply GITHUB_TOKEN with Copilot access when recording new provider responses. +cd nodejs +npm test -- test/e2e/structured_output.e2e.test.ts +cd ../dotnet +dotnet test test/GitHub.Copilot.SDK.Test.csproj \ + --filter FullyQualifiedName~StructuredOutputE2ETests +``` + +The shared harness records real inference responses under `test/snapshots`. +Record new captures with `GITHUB_TOKEN` set and `GITHUB_ACTIONS` unset; +never author model responses by hand. Rerun with `GITHUB_ACTIONS=true` and real +provider credentials removed to require replay instead of forwarding cache +misses upstream. A draft targeting an unreleased runtime should document the +required runtime revision; update the pinned release only after it ships. +Pinned-schema CI can report drift in such a draft, and Java codegen may +automatically update generated files to match the pinned release. + +For recording behind `HTTPS_PROXY`, Node versions that support environment +proxies (including Node 24.20) need `NODE_USE_ENV_PROXY=1` in the test runner's +environment. If the host proxy substitutes a protected credential, set +`GITHUB_TOKEN="$GH_TOKEN"` using its issued placeholder; do not print or persist +the credential. Keep localhost and loopback in `NO_PROXY`. + +Equivalent cross-language E2Es must share snapshot names and prompts, not +language-specific copies. The structured-output suite in **all six SDKs** reuses +the following captures in `test/snapshots/structured_output/`, recorded using +real CAPI `gpt-4.1` calls through the shared harness: + +| Shared capture (without `.yaml`) | Flow | +| --- | --- | +| `infers_typed_result_after_custom_tool` | Inferred typed result after a tool call, streamed text, then an unformatted follow-up | +| `sends_explicit_schema_for_message_and_batch` | Explicit-schema batch RPC followed by a schema-bearing single send | +| `send_selects_correlated_response_after_idle` | Event-driven send, tool commentary, originating-message correlation, and an idle boundary held by a stop hook | +| `typed_wait_returns_stop_hook_correction` | Typed wait returns the corrected answer, not the first assistant message | +| `typed_wait_returns_stop_hook_correction_after_terminal_tool` | Output-only finalization after a terminal tool, followed by a stop-hook correction | +| `typed_result_after_terminal_tool_and_steering` | Immediate steering during a terminal tool preserves the active schema | +| `typed_wait_returns_late_steering_response` | Steering after the first final answer remains part of the original run | +| `concurrent_typed_sends_return_their_own_results` | Concurrent queued runs use different inferred types and return their own results | + +Typed cases call the public idiomatic APIs: Node/Zod, C# generics, Python/Pydantic, +Go generics, Java annotated records using the existing tool schema generator, +and Rust generics with `derive`/schemars. The tool/follow-up case also checks the +actual provider request's inferred schema, so a recorded JSON response alone +cannot mask missing schema forwarding. Explicit-schema and event-stream cases +exercise the corresponding raw public APIs instead. + +Every language additionally checks rejection before admission and zero provider +calls for oversized schemas and typed immediate steering. These cases have no +model responses and therefore need **no snapshot**. Do not create canned responses +or empty model captures for them. Unit tests supplement, rather than replace, +the shared runtime E2Es. + ## Submitting a Pull Request 1. Fork and clone the repository diff --git a/dotnet/README.md b/dotnet/README.md index f8b7e44f06..0e764b6c2d 100644 --- a/dotnet/README.md +++ b/dotnet/README.md @@ -251,6 +251,7 @@ Send a message to the session. - `Attachments` - File attachments - `Mode` - Delivery mode ("enqueue" or "immediate") - `Source` - Optional message origin: `MessageSource.User`, `MessageSource.System`, or `MessageSource.Agent(id)`. Omitted by default, preserving the runtime's default user behavior. +- `ResponseSchema` - Experimental provider-native JSON Schema (`JsonElement`) for this turn. Returns the message ID. @@ -277,6 +278,132 @@ await session.SendAndWaitAsync(new MessageOptions Agent sources serialize as `agent-`. Pass the agent ID without adding a prefix. The SDK preserves its case and whitespace and rejects null IDs. +##### Structured outputs (experimental) + +Use `SendAndWaitAsync` to infer a JSON Schema from a .NET type and +deserialize the final response. Schema inference uses +`Microsoft.Extensions.AI.AIJsonUtilities`, the same technology as custom tools. +In a reflection-enabled application, `await session.SendAndWaitAsync(prompt)` +needs no serialization configuration. The example below supplies source-generated +metadata so it also works when reflection serialization is disabled. + +```csharp +var result = await session.SendAndWaitAsync( + "How many red widgets are in stock?", + serializerOptions: InventoryJsonContext.Default.Options); +Console.WriteLine($"{result.Count} {result.Color} widgets"); + +public sealed class Inventory +{ + public required int Count { get; set; } + public required string Color { get; set; } +} + +[System.Text.Json.Serialization.JsonSourceGenerationOptions( + PropertyNamingPolicy = System.Text.Json.Serialization.JsonKnownNamingPolicy.CamelCase)] +[System.Text.Json.Serialization.JsonSerializable(typeof(Inventory))] +internal partial class InventoryJsonContext : System.Text.Json.Serialization.JsonSerializerContext; +``` + +The same serialization options govern schema inference and deserialization, +including naming policies, `[JsonPropertyName]`, converters, required members, +and nullable annotations. Options default to `AIJsonUtilities.DefaultOptions`, +as for custom tools. Supply a source-generated resolver (as above) for Native +AOT or when reflection serialization is disabled. The typed helper requests +strict output, marks all schema properties required, and disallows additional +properties; nullable properties can still contain JSON null. + +The helper waits for non-autopilot session idle after the requested user message +is consumed, selecting only root assistant messages with that originating message +ID. This can wait for other queued work to drain, but other messages and subagent +responses cannot replace the result. Session errors or an aborted idle after the +requested run starts conservatively fail the wait, even if later queued work +caused them. It throws `InvalidOperationException` when there is no final response, +and `JsonException` for invalid JSON, an incompatible +value, or a null result. Deserialization is not full JSON Schema validation: +validate application-specific constraints yourself. Timeout defaults to 60 +seconds; timeout and cancellation stop waiting without aborting runtime work. +The original `MessageOptions` is not modified, and an explicit `ResponseSchema` +cannot be combined with this typed overload. + +For an explicit schema, set `MessageOptions.ResponseSchema`. Schemas are opaque +`JsonElement` values, just like custom-tool schemas. The SDK forwards this schema +unchanged with the name `response` and `strict: true`. The untyped +`SendAndWaitAsync` still returns an assistant message event; it does not validate +or deserialize the response. Schema-bearing waits use the same message +correlation as typed waits; unformatted waits retain their existing behavior. + +With `SendAsync`, collect root `AssistantMessageEvent` events whose +`Data.OriginatingMessageId` matches the returned message ID, then select the last +one without tool requests when the session becomes idle. Subscribe before sending +because events can precede the send acknowledgement, and handle `SessionErrorEvent` normally. +There is no final-message flag: stop hooks can reject an initial answer and +request a correction. Those corrections retain the original schema and +originating message ID, so `SendAndWaitAsync` selects the corrected response at +idle. Independent queued sends retain their own schemas and IDs. + +```csharp +using var schema = System.Text.Json.JsonDocument.Parse(""" + {"type":"object","properties":{"count":{"type":"integer"}},"required":["count"],"additionalProperties":false} + """); +var message = await session.SendAndWaitAsync(new MessageOptions +{ + Prompt = "Count the widgets.", + ResponseSchema = schema.RootElement.Clone(), +}); +``` + +Use the generated `session.Rpc` APIs for advanced response-format options: + +```csharp +using GitHub.Copilot.Rpc; +using System.Text.Json; + +using var schema = JsonDocument.Parse(""" + {"type":"object","properties":{"count":{"type":"integer"}},"required":["count"],"additionalProperties":false} + """); +var format = new ResponseFormatJsonSchema +{ + JsonSchema = new JsonSchemaResponseFormat + { + Name = "inventory", + Schema = schema.RootElement.Clone(), + Strict = true, + Description = "The inventory count", + }, +}; +await session.Rpc.SendAsync("Count the widgets.", responseFormat: format); +// A batch shares one output contract: +await session.Rpc.SendMessagesAsync( + [new() { Prompt = "There are 42 widgets." }, new() { Prompt = "Report the count." }], + responseFormat: format); +``` + +Raw schemas and outputs are passed through without validation or rewriting. +Provider support and schema restrictions apply. The format persists through +tool continuations in that run, not independent subsequent runs. An ordinary +`Mode = "immediate"` steering message inherits the active format and originating +message ID, even if it arrives after the final model request and is promoted +into a follow-up run. Specifying a new format on an immediate message is rejected, +even while idle. +Each batch starts one run: the final returned message ID is its origin, preceding +messages are context, and an empty batch has no origin. An immediate batch +steers the active run instead and retains its origin. +The schema is not a persisted session default: autonomous resume-pending work +after a restart does not restore it. A terminal tool that clears context ends +the old run; its fresh seed does not inherit the schema or origin. Such a run +can finish without a structured result, in which case the typed wait throws. +After a successful terminal tool, the runtime disables tools while the model +produces the structured result. Stop-hook corrections remain supported. +Remote sessions and known HydraFusion routes reject response formats before +admission. Schemas larger than 32 MiB when JSON-encoded are also rejected before +admission, using the runtime's existing request-size ceiling. This does not +guarantee the schema plus conversation and tools fits the provider's budget. +Use a provider route that enforces JSON Schema: an API-compatible gateway can +ignore unsupported format fields, and the Claude Chat-completions compatibility +route is not equivalent to Anthropic's native Messages endpoint. The SDK's +pinned CLI release includes the required runtime support. + ##### `On(Action handler): IDisposable` Subscribe to session events. Returns a disposable to unsubscribe. diff --git a/dotnet/src/Session.StructuredOutput.cs b/dotnet/src/Session.StructuredOutput.cs new file mode 100644 index 0000000000..6df54d121e --- /dev/null +++ b/dotnet/src/Session.StructuredOutput.cs @@ -0,0 +1,207 @@ +/*--------------------------------------------------------------------------------------------- + * Copyright (c) Microsoft Corporation. All rights reserved. + *--------------------------------------------------------------------------------------------*/ + +using Microsoft.Extensions.AI; +using System.Diagnostics.CodeAnalysis; +using System.Text.Json; +using System.Text.Json.Serialization.Metadata; + +namespace GitHub.Copilot; + +public sealed partial class CopilotSession +{ + /// + /// Sends a prompt with a JSON Schema inferred from and + /// deserializes the final response into that type. + /// + /// The expected response type. + /// The user message text. + /// Options used both for schema inference and deserialization. + /// Defaults to , as for custom tools. + /// For Native AOT, supply options with a source-generated type resolver. + /// Timeout duration (default: 60 seconds). Does not abort agent work. + /// Cancellation token for sending and waiting. + /// The non-null deserialized response. + [Experimental(Diagnostics.Experimental)] + public Task SendAndWaitAsync( + string prompt, + JsonSerializerOptions? serializerOptions = null, + TimeSpan? timeout = null, + CancellationToken cancellationToken = default) + { + ArgumentNullException.ThrowIfNull(prompt); + return SendAndWaitAsync(new MessageOptions { Prompt = prompt }, serializerOptions, timeout, cancellationToken); + } + + /// + /// Sends a message with a JSON Schema inferred from and + /// deserializes the final response into that type. + /// + /// The expected response type. + /// The message to send. Must not specify a response schema or immediate delivery. + /// Options used both for schema inference and deserialization. + /// Defaults to , as for custom tools. + /// For Native AOT, supply options with a source-generated type resolver. + /// Timeout duration (default: 60 seconds). Does not abort agent work. + /// Cancellation token for sending and waiting. + /// The non-null deserialized response. + /// The message specifies a response schema or immediate delivery. + /// No final response was received, or the session reported an error. + /// The response is not valid JSON for the requested type, or is null. + /// The response did not arrive within the timeout. + /// + /// Uses the same Microsoft.Extensions.AI schema inference as custom tools. Property naming, + /// converters, required members and nullable annotations follow the supplied serialization + /// contracts. The inferred schema requests strict output with all properties required and + /// additional properties disallowed. Provider schema restrictions still apply. + /// Deserialization is not full JSON Schema validation; apply application-specific validation + /// to the returned value where needed. The supplied message options are not modified. + /// + [Experimental(Diagnostics.Experimental)] + public async Task SendAndWaitAsync( + MessageOptions options, + JsonSerializerOptions? serializerOptions = null, + TimeSpan? timeout = null, + CancellationToken cancellationToken = default) + { + ArgumentNullException.ThrowIfNull(options); + ThrowIfDisposed(); + if (options.ResponseSchema is not null) + { + throw new ArgumentException("The typed overload infers its response schema. Use the untyped overload for an explicit response schema.", nameof(options)); + } + if (options.Mode == "immediate") + { + throw new ArgumentException("Structured output cannot be requested on an immediate steering message.", nameof(options)); + } + + serializerOptions = ResolveStructuredOutputOptions(serializerOptions); + var typeInfo = (JsonTypeInfo)serializerOptions.GetTypeInfo(typeof(TResult)); + var schema = AIJsonUtilities.CreateJsonSchema( + typeof(TResult), + serializerOptions: serializerOptions, + inferenceOptions: new AIJsonSchemaCreateOptions + { + TransformOptions = new AIJsonSchemaTransformOptions + { + RequireAllProperties = true, + DisallowAdditionalProperties = true, + MoveDefaultKeywordToDescription = true, + }, + }); + var message = options.Clone(); + message.ResponseSchema = schema; + + var response = await SendAndWaitForStructuredMessageAsync(message, timeout, cancellationToken); + return JsonSerializer.Deserialize(response.Data.Content, typeInfo) + ?? throw new JsonException("The structured response was JSON null, not a result."); + } + + [UnconditionalSuppressMessage("AOT", "IL3050", Justification = "The reflection resolver is only created when JsonSerializer.IsReflectionEnabledByDefault is enabled.")] + [UnconditionalSuppressMessage("Trimming", "IL2026", Justification = "The reflection resolver is only created when JsonSerializer.IsReflectionEnabledByDefault is enabled.")] + private static JsonSerializerOptions ResolveStructuredOutputOptions(JsonSerializerOptions? options) + { + options ??= AIJsonUtilities.DefaultOptions; + if (options.IsReadOnly) + { + return options; + } + var resolved = new JsonSerializerOptions(options); + if (resolved.TypeInfoResolver is null && JsonSerializer.IsReflectionEnabledByDefault) + { + resolved.TypeInfoResolver = new DefaultJsonTypeInfoResolver(); + } + resolved.MakeReadOnly(); + return resolved; + } + + private async Task SendAndWaitForStructuredMessageAsync( + MessageOptions options, TimeSpan? timeout, CancellationToken cancellationToken) + { + var effectiveTimeout = timeout ?? TimeSpan.FromSeconds(60); + using var cts = CancellationTokenSource.CreateLinkedTokenSource(cancellationToken); + cts.CancelAfter(effectiveTimeout); + var completion = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + using var registration = cts.Token.Register(() => completion.TrySetCanceled(cts.Token)); + var gate = new object(); + var pendingEvents = new List(); + string? messageId = null; + var started = false; + AssistantMessageEvent? finalMessage = null; + + void ProcessEvent(SessionEvent evt) + { + switch (evt) + { + case UserMessageEvent user when string.IsNullOrEmpty(user.AgentId) && user.Data.MessageId == messageId: + started = true; + break; + case AssistantMessageEvent assistant when string.IsNullOrEmpty(assistant.AgentId) && assistant.Data.OriginatingMessageId == messageId: + started = true; + finalMessage = assistant.Data.ToolRequests is { Length: > 0 } ? null : assistant; + break; + case SessionIdleEvent idle when started && string.IsNullOrEmpty(idle.AgentId) && idle.Data.Mode != SessionMode.Autopilot: + if (idle.Data.Aborted == true) + { + completion.TrySetException(new InvalidOperationException("The session was aborted before a final structured response was received.")); + } + else if (finalMessage is null || string.IsNullOrWhiteSpace(finalMessage.Data.Content)) + { + completion.TrySetException(new InvalidOperationException("The turn completed without a final structured response.")); + } + else + { + completion.TrySetResult(finalMessage); + } + break; + case SessionErrorEvent error when started && string.IsNullOrEmpty(error.AgentId): + completion.TrySetException(new InvalidOperationException($"Session error: {error.Data.Message}")); + break; + } + } + + using var subscription = On(evt => + { + if (evt is not (UserMessageEvent or AssistantMessageEvent or SessionIdleEvent or SessionErrorEvent)) + { + return; + } + lock (gate) + { + if (messageId is null) + { + // Events can arrive before the send RPC response supplies the logical message ID. + pendingEvents.Add(evt); + } + else + { + ProcessEvent(evt); + } + } + }); + try + { + var sentMessageId = await SendAsync(options, cts.Token); + lock (gate) + { + messageId = sentMessageId; + foreach (var evt in pendingEvents) + { + ProcessEvent(evt); + } + pendingEvents.Clear(); + } + await Task.WhenAny(completion.Task, JsonRpc.Completion, _eventChannel.Reader.Completion); + if (!completion.Task.IsCompleted) + { + throw new IOException("The session closed before a final structured response was received."); + } + return await completion.Task; + } + catch (OperationCanceledException) when (!cancellationToken.IsCancellationRequested) + { + throw new TimeoutException($"SendAndWaitAsync timed out after {effectiveTimeout}"); + } + } +} diff --git a/dotnet/src/Session.cs b/dotnet/src/Session.cs index 404c0054b7..3b2e1c349d 100644 --- a/dotnet/src/Session.cs +++ b/dotnet/src/Session.cs @@ -289,7 +289,8 @@ public Task SendAsync(string prompt, CancellationToken cancellationToken /// /// Options for the message to be sent, including the prompt and optional attachments. /// A that can be used to cancel the operation. - /// A task that resolves with the ID of the response message, which can be used to correlate events. + /// The submitted user message's ID, not an assistant response ID. When this send starts + /// a run, root assistant messages carry it as OriginatingMessageId. /// Thrown if the session has been disposed. /// /// @@ -331,6 +332,15 @@ public async Task SendAsync(MessageOptions options, CancellationToken ca Traceparent = traceparent, Tracestate = tracestate, RequestHeaders = options.RequestHeaders, + ResponseFormat = options.ResponseSchema is { } schema ? new ResponseFormatJsonSchema + { + JsonSchema = new JsonSchemaResponseFormat + { + Name = "response", + Schema = schema, + Strict = true, + }, + } : null, }; var rpcTimestamp = Stopwatch.GetTimestamp(); @@ -380,6 +390,11 @@ public async Task SendAsync(MessageOptions options, CancellationToken ca ArgumentNullException.ThrowIfNull(options); ThrowIfDisposed(); + if (options.ResponseSchema is not null) + { + return await SendAndWaitForStructuredMessageAsync(options, timeout, cancellationToken); + } + var totalTimestamp = Stopwatch.GetTimestamp(); var effectiveTimeout = timeout ?? TimeSpan.FromSeconds(60); var tcs = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); @@ -2275,6 +2290,7 @@ internal record SendMessageRequest public string? Traceparent { get; init; } public string? Tracestate { get; init; } public IDictionary? RequestHeaders { get; init; } + public ResponseFormat? ResponseFormat { get; init; } } internal record SendMessageResponse diff --git a/dotnet/src/Types.cs b/dotnet/src/Types.cs index 1cc4919093..4cda151f1d 100644 --- a/dotnet/src/Types.cs +++ b/dotnet/src/Types.cs @@ -4091,6 +4091,7 @@ private MessageOptions(MessageOptions? other) Source = other.Source; Prompt = other.Prompt; DisplayPrompt = other.DisplayPrompt; + ResponseSchema = other.ResponseSchema; RequestHeaders = other.RequestHeaders is not null ? new Dictionary(other.RequestHeaders) : null; @@ -4128,6 +4129,18 @@ private MessageOptions(MessageOptions? other) /// public string? DisplayPrompt { get; set; } + /// + /// Optional provider-native JSON Schema for this turn, including tool continuations. + /// The schema is passed unchanged with the name "response" and strict enforcement requested. + /// Ordinary immediate steering retains the active run's schema and origin even when promoted + /// to a follow-up after the model request finishes. An immediate message must not specify its + /// own schema, even while idle. Independent sends and context resets do not inherit this schema; + /// it is not a persisted session default. + /// Use for advanced response-format options. + /// + [Experimental(Diagnostics.Experimental)] + public JsonElement? ResponseSchema { get; set; } + /// /// Creates a shallow clone of this instance. /// diff --git a/dotnet/test/E2E/StructuredOutputE2ETests.cs b/dotnet/test/E2E/StructuredOutputE2ETests.cs new file mode 100644 index 0000000000..e09a796a9d --- /dev/null +++ b/dotnet/test/E2E/StructuredOutputE2ETests.cs @@ -0,0 +1,506 @@ +/*--------------------------------------------------------------------------------------------- + * Copyright (c) Microsoft Corporation. All rights reserved. + *--------------------------------------------------------------------------------------------*/ + +using GitHub.Copilot.Rpc; +using GitHub.Copilot.Test.Harness; +using System.Text.Json; +using System.Text.Json.Serialization; +using Xunit; +using Xunit.Abstractions; + +namespace GitHub.Copilot.Test.E2E; + +public partial class StructuredOutputE2ETests(E2ETestFixture fixture, ITestOutputHelper output) + : E2ETestBase(fixture, "structured_output", output) +{ + private SessionConfig StructuredSessionConfig() => E2ETestBackendConfiguration.Current != E2ETestBackend.Capi + ? new() { AvailableTools = [] } + : new() + { + Model = "gpt-4.1", + AvailableTools = [], + Provider = new ProviderConfig + { + Type = "openai", + WireApi = "completions", + BaseUrl = Ctx.ProxyUrl, + ModelId = "gpt-4.1", + WireModel = "gpt-4.1", + ApiKey = Environment.GetEnvironmentVariable("GITHUB_ACTIONS") == "true" + ? "fake-token-for-e2e-tests" + : Environment.GetEnvironmentVariable("GITHUB_TOKEN") ?? "fake-token-for-e2e-tests", + Headers = new Dictionary + { + ["Copilot-Integration-Id"] = "copilot-developer-cli", + ["Copilot-Harness-Id"] = "copilot-sdk", + ["X-GitHub-Api-Version"] = "2026-08-01", + }, + }, + }; + + [Fact] + public async Task Infers_Typed_Result_After_Custom_Tool() + { + var calls = 0; + var config = StructuredSessionConfig(); + config.Streaming = true; + config.Tools = + [ + CopilotTool.DefineTool(() => + { + calls++; + return "The inventory contains 42 red widgets."; + }, factoryOptions: new() { Name = "get_inventory", Description = "Get the current widget inventory." }), + ]; + var session = await CreateSessionAsync(config); + var deltas = 0; + using var subscription = session.On(_ => Interlocked.Increment(ref deltas)); + + var serializerOptions = new JsonSerializerOptions { PropertyNamingPolicy = JsonNamingPolicy.CamelCase }; + if (!JsonSerializer.IsReflectionEnabledByDefault) + { + serializerOptions.TypeInfoResolver = StructuredOutputE2EJsonContext.Default; + } + var result = await session.SendAndWaitAsync( + "Call get_inventory, then report the widget count and color.", + serializerOptions, + TimeSpan.FromMinutes(3)); + Assert.False(serializerOptions.IsReadOnly); + Assert.True(calls > 0); + Assert.True(deltas > 0, "Typed wait must preserve streaming text updates"); + Assert.Equal(42, result.Count); + Assert.Equal("red", result.Color); + + var ordinary = await session.SendAndWaitAsync( + "Now reply with exactly the plain text HELLO, not JSON.", + TimeSpan.FromMinutes(3)); + Assert.NotNull(ordinary); + Assert.Equal("HELLO", ordinary.Data.Content.Trim()); + if (E2ETestBackendConfiguration.Current == E2ETestBackend.Capi) + { + var exchanges = await Ctx.GetExchangesAsync(); + Assert.True(exchanges.Count >= 3); + foreach (var exchange in exchanges.Take(exchanges.Count - 1)) + { + Assert.NotNull(exchange.Request.ResponseFormat); + var format = exchange.Request.ResponseFormat.Value; + Assert.Equal("json_schema", format.GetProperty("type").GetString()); + var contract = format.GetProperty("json_schema"); + Assert.True(contract.GetProperty("strict").GetBoolean()); + var schema = contract.GetProperty("schema"); + Assert.False(schema.GetProperty("additionalProperties").GetBoolean()); + Assert.Equal("integer", schema.GetProperty("properties").GetProperty("count").GetProperty("type").GetString()); + Assert.Equal("string", schema.GetProperty("properties").GetProperty("color").GetProperty("type").GetString()); + } + Assert.Null(exchanges.Last().Request.ResponseFormat); + } + } + + [Fact] + public async Task Sends_Explicit_Schema_For_Message_And_Batch() + { + var session = await CreateSessionAsync(StructuredSessionConfig()); + var completion = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + var replies = new System.Collections.Concurrent.ConcurrentQueue(); + using var subscription = session.On(evt => + { + if (!string.IsNullOrEmpty(evt.AgentId)) return; + switch (evt) + { + case AssistantMessageEvent message: + replies.Enqueue(message); + break; + case SessionIdleEvent: + completion.TrySetResult(); + break; + case SessionErrorEvent error: + completion.TrySetException(new InvalidOperationException(error.Data.Message)); + break; + } + }); + using var schema = JsonDocument.Parse( + """{"type":"object","properties":{"count":{"type":"integer"},"color":{"type":"string"}},"required":["count","color"],"additionalProperties":false}"""); + var accepted = await session.Rpc.SendMessagesAsync( + [new() { Prompt = "There are 42 red widgets in stock." }, new() { Prompt = "Report the widget count and color." }], + responseFormat: new ResponseFormatJsonSchema + { + JsonSchema = new JsonSchemaResponseFormat + { + Name = "inventory", + Schema = schema.RootElement.Clone(), + Strict = true, + Description = "The widget inventory", + }, + }); + using var cts = new CancellationTokenSource(TimeSpan.FromMinutes(3)); + await completion.Task.WaitAsync(Timeout.InfiniteTimeSpan, cts.Token); + var message = replies.Last(message => message.Data.OriginatingMessageId == accepted.MessageIds.Last()); + Assert.Equal(accepted.MessageIds.Last(), message.Data.OriginatingMessageId); + Assert.Empty(message.Data.ToolRequests ?? []); + var result = JsonSerializer.Deserialize(message.Data.Content, StructuredOutputE2EJsonContext.Default.Inventory); + Assert.NotNull(result); + Assert.Equal(42, result.Count); + Assert.Equal("red", result.Color); + + var raw = await session.SendAndWaitAsync(new MessageOptions + { + Prompt = "The inventory now has 21 blue widgets. Report the new count and color.", + ResponseSchema = schema.RootElement.Clone(), + }, TimeSpan.FromMinutes(3)); + Assert.NotNull(raw); + var updated = JsonSerializer.Deserialize(raw.Data.Content, StructuredOutputE2EJsonContext.Default.Inventory); + Assert.NotNull(updated); + Assert.Equal(21, updated.Count); + Assert.Equal("blue", updated.Color); + } + + [Fact] + public async Task Send_Selects_Correlated_Response_After_Idle() + { + var hookEntered = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + var releaseHook = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + var config = StructuredSessionConfig(); + config.Tools = + [ + CopilotTool.DefineTool(() => "The inventory contains 42 red widgets.", + factoryOptions: new() { Name = "read_inventory", Description = "Read the current widget count and color." }), + ]; + config.Hooks = new SessionHooks + { + OnAgentStop = async (_, _) => + { + hookEntered.TrySetResult(); + await releaseHook.Task; + return null; + }, + }; + var session = await CreateSessionAsync(config); + var idleReceived = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + var replies = new System.Collections.Concurrent.ConcurrentQueue(); + using var subscription = session.On(evt => + { + if (!string.IsNullOrEmpty(evt.AgentId)) return; + switch (evt) + { + case AssistantMessageEvent message: + replies.Enqueue(message); + break; + case SessionErrorEvent error: + idleReceived.TrySetException(new InvalidOperationException(error.Data.Message)); + break; + case SessionIdleEvent: + idleReceived.TrySetResult(); + break; + } + }); + using var cts = new CancellationTokenSource(TimeSpan.FromMinutes(3)); + using var schema = JsonDocument.Parse( + """{"type":"object","properties":{"count":{"type":"integer"},"color":{"type":"string"}},"required":["count","color"],"additionalProperties":false}"""); + try + { + var messageId = await session.SendAsync(new MessageOptions + { + Prompt = "Call read_inventory once, then report the current widget count and color.", + ResponseSchema = schema.RootElement.Clone(), + }, cts.Token); + await hookEntered.Task.WaitAsync(Timeout.InfiniteTimeSpan, cts.Token); + Assert.False(idleReceived.Task.IsCompleted); + releaseHook.TrySetResult(); + await idleReceived.Task.WaitAsync(Timeout.InfiniteTimeSpan, cts.Token); + var reply = replies.Last(message => message.Data.OriginatingMessageId == messageId); + Assert.Equal(messageId, reply.Data.OriginatingMessageId); + var result = JsonSerializer.Deserialize(reply.Data.Content, StructuredOutputE2EJsonContext.Default.Inventory); + Assert.NotNull(result); + Assert.Equal(42, result.Count); + Assert.Equal("red", result.Color); + Assert.Contains(replies, message => message.Data.ToolRequests is { Length: > 0 }); + Assert.Empty(reply.Data.ToolRequests ?? []); + Assert.Same(reply, replies.Last()); + } + finally + { + releaseHook.TrySetResult(); + } + } + + [Fact] + public async Task Typed_Wait_Returns_Stop_Hook_Correction() + { + var stops = 0; + var config = StructuredSessionConfig(); + config.Hooks = new SessionHooks + { + OnAgentStop = (_, _) => Task.FromResult( + Interlocked.Increment(ref stops) == 1 + ? new() { Decision = "block", Reason = "Correct the answer to 99, not 42. Do not use tools." } + : null), + }; + var session = await CreateSessionAsync(config); + var replies = new System.Collections.Concurrent.ConcurrentQueue(); + using var subscription = session.On(message => + { + if (string.IsNullOrEmpty(message.AgentId)) replies.Enqueue(message); + }); + var result = await session.SendAndWaitAsync( + "What is 19 + 23? Do not use tools.", + StructuredOutputE2EJsonContext.Default.Options, + TimeSpan.FromMinutes(3)); + Assert.Equal(99, result.Answer); + Assert.Equal(2, stops); + Assert.Equal(2, replies.Count); + Assert.False(string.IsNullOrEmpty(replies.First().Data.OriginatingMessageId)); + Assert.Equal(replies.First().Data.OriginatingMessageId, replies.Last().Data.OriginatingMessageId); + Assert.Equal([42, 99], replies.Select(message => + JsonSerializer.Deserialize(message.Data.Content, StructuredOutputE2EJsonContext.Default.CorrectionResult)!.Answer)); + } + + [Fact] + public async Task Typed_Wait_Returns_Late_Steering_Response() + { + var stops = 0; + string? steeringId = null; + CopilotSession? session = null; + var config = StructuredSessionConfig(); + config.Hooks = new SessionHooks + { + OnAgentStop = async (_, _) => + { + if (Interlocked.Increment(ref stops) == 1) + { + // The final model request has finished, but this run still admits steering. + steeringId = await session!.SendAsync(new MessageOptions + { + Prompt = "Change the answer to 99. Do not use tools.", + Mode = "immediate", + }); + } + return null; + }, + }; + session = await CreateSessionAsync(config); + var replies = new System.Collections.Concurrent.ConcurrentQueue(); + using var subscription = session.On(message => + { + if (string.IsNullOrEmpty(message.AgentId)) replies.Enqueue(message); + }); + var result = await session.SendAndWaitAsync( + "What is 19 + 23? Do not use tools.", + StructuredOutputE2EJsonContext.Default.Options, + TimeSpan.FromMinutes(3)); + Assert.Equal(99, result.Answer); + Assert.Equal(2, stops); + Assert.Equal(2, replies.Count); + Assert.False(string.IsNullOrEmpty(steeringId)); + Assert.False(string.IsNullOrEmpty(replies.First().Data.OriginatingMessageId)); + Assert.NotEqual(steeringId, replies.First().Data.OriginatingMessageId); + Assert.Equal(replies.First().Data.OriginatingMessageId, replies.Last().Data.OriginatingMessageId); + Assert.Equal([42, 99], replies.Select(message => + JsonSerializer.Deserialize(message.Data.Content, StructuredOutputE2EJsonContext.Default.CorrectionResult)!.Answer)); + } + + [Fact] + public async Task Typed_Wait_Returns_Stop_Hook_Correction_After_Terminal_Tool() + { + var calls = 0; + var stops = 0; + var config = StructuredSessionConfig(); + config.Tools = + [ + CopilotTool.DefineTool(() => + { + Interlocked.Increment(ref calls); + return 58; + }, new CopilotToolOptions { IsTerminal = true, SkipPermission = true }, + new() { Name = "lookup_number", Description = "Return the number needed for the calculation." }), + ]; + config.Hooks = new SessionHooks + { + OnAgentStop = (_, _) => Task.FromResult( + Interlocked.Increment(ref stops) == 1 + ? new() { Decision = "block", Reason = "Correct the answer to 99, not 63. Do not use tools." } + : null), + }; + var session = await CreateSessionAsync(config); + var replies = new System.Collections.Concurrent.ConcurrentQueue(); + using var subscription = session.On(message => + { + if (string.IsNullOrEmpty(message.AgentId)) replies.Enqueue(message); + }); + var result = await session.SendAndWaitAsync( + "Call lookup_number exactly once, then add 5 to the returned number. Do not guess its result.", + StructuredOutputE2EJsonContext.Default.Options, + TimeSpan.FromMinutes(3)); + Assert.Equal(99, result.Answer); + Assert.Equal(1, calls); + Assert.Equal(2, stops); + var answers = replies.Where(message => message.Data.ToolRequests is not { Length: > 0 }).ToArray(); + Assert.Equal([63, 99], answers.Select(message => + JsonSerializer.Deserialize(message.Data.Content, StructuredOutputE2EJsonContext.Default.CorrectionResult)!.Answer)); + Assert.False(string.IsNullOrEmpty(answers[0].Data.OriginatingMessageId)); + Assert.Equal(answers[0].Data.OriginatingMessageId, answers[1].Data.OriginatingMessageId); + var exchanges = await Ctx.GetExchangesAsync(); + Assert.Equal(3, exchanges.Count); + Assert.Equal("none", exchanges[1].Request.ToolChoice?.GetString()); + } + + [Fact] + public async Task Typed_Result_After_Terminal_Tool_And_Steering() + { + var calls = 0; + CopilotSession? session = null; + var config = StructuredSessionConfig(); + config.Tools = + [ + CopilotTool.DefineTool(async () => + { + Interlocked.Increment(ref calls); + await session!.SendAsync(new MessageOptions + { + Prompt = "Continue with the original calculation. Do not call any more tools.", + Mode = "immediate", + }); + return 58; + }, new CopilotToolOptions { IsTerminal = true, SkipPermission = true }, + new() { Name = "lookup_number", Description = "Return the number needed for the calculation." }), + ]; + session = await CreateSessionAsync(config); + var result = await session.SendAndWaitAsync( + "Call lookup_number exactly once, then add 5 to the returned number. Do not guess its result.", + StructuredOutputE2EJsonContext.Default.Options, + TimeSpan.FromMinutes(3)); + Assert.Equal(63, result.Answer); + Assert.Equal("typed_tool", result.Contract); + Assert.Equal(1, calls); + var exchanges = await Ctx.GetExchangesAsync(); + Assert.True(exchanges.Count >= 2); + Assert.All(exchanges.Skip(1), exchange => Assert.Equal("none", exchange.Request.ToolChoice?.GetString())); + if (E2ETestBackendConfiguration.Current == E2ETestBackend.Capi) + { + Assert.All(exchanges, exchange => Assert.Equal("json_schema", + exchange.Request.ResponseFormat?.GetProperty("type").GetString())); + } + } + + [Fact] + public async Task Rejects_Unsupported_Or_Oversized_Schemas_Before_Admission() + { + var environment = Ctx.GetEnvironment(); + environment["COPILOT_CLI_ENABLED_FEATURE_FLAGS"] = "HYDRAFUSION,HYDRAFUSION_ROLLOUT"; + await using var client = Ctx.CreateClient(environment: environment); + foreach (var model in new[] { "gpt-4.1", "hydrafusion" }) + { + var config = StructuredSessionConfig(); + config.Model = model; + config.OnPermissionRequest = PermissionHandler.ApproveAll; + await using var session = await Ctx.CreateSessionAsync(client, config); + await Assert.ThrowsAsync(() => session.SendAndWaitAsync( + new MessageOptions { Prompt = "Must not be admitted", Mode = "immediate" }, + StructuredOutputE2EJsonContext.Default.Options)); + using var schema = JsonDocument.Parse( + "{\"type\":\"object\",\"description\":\"" + + (model == "gpt-4.1" ? new string('x', 32 * 1024 * 1024) : "Small schema") + "\"}"); + var message = model == "gpt-4.1" ? "32 MiB" : "HydraFusion"; + var error = await Assert.ThrowsAnyAsync(() => + session.SendAndWaitAsync(new MessageOptions + { + Prompt = "Must not be admitted", + ResponseSchema = schema.RootElement, + })); + Assert.Contains(message, error.Message); + error = await Assert.ThrowsAnyAsync(() => + session.Rpc.SendMessagesAsync([], responseFormat: new ResponseFormatJsonSchema + { + JsonSchema = new() { Name = "response", Schema = schema.RootElement }, + })); + Assert.Contains(message, error.Message); + Assert.Empty((await session.Rpc.Queue.PendingItemsAsync()).Items); + Assert.DoesNotContain(await session.GetEventsAsync(), + evt => evt is UserMessageEvent or SessionErrorEvent); + } + Assert.Empty(await Ctx.GetExchangesAsync()); + } + + [Fact] + public async Task Concurrent_Typed_Sends_Return_Their_Own_Results() + { + var toolEntered = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + var releaseTool = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + var config = StructuredSessionConfig(); + config.Tools = + [ + CopilotTool.DefineTool(async () => + { + toolEntered.TrySetResult(); + await releaseTool.Task; + return 42; + }, factoryOptions: new() { Name = "first_number", Description = "Get the number for the first question." }), + ]; + var session = await CreateSessionAsync(config); + var serializerOptions = JsonSerializer.IsReflectionEnabledByDefault + ? null + : StructuredOutputE2EJsonContext.Default.Options; + var first = session.SendAndWaitAsync( + "Call first_number exactly once and report its returned number.", + serializerOptions, + TimeSpan.FromMinutes(3)); + try + { + var entered = await Task.WhenAny(toolEntered.Task, first).WaitAsync(TimeSpan.FromMinutes(3)); + if (entered == first) + { + await first; + throw new InvalidOperationException("First run completed without calling first_number."); + } + const string secondPrompt = "What is 30 + 7? Do not use tools."; + var second = session.SendAndWaitAsync( + secondPrompt, serializerOptions, TimeSpan.FromMinutes(3)); + await TestHelper.WaitForConditionAsync( + async () => (await session.Rpc.Queue.PendingItemsAsync()).Items.Any( + item => item.DisplayText.Contains(secondPrompt, StringComparison.Ordinal)), + timeoutMessage: "Second structured send was not queued behind the tool call."); + releaseTool.TrySetResult(); + Assert.Equal(42, (await first).First); + Assert.Equal(37, (await second).Second); + } + finally + { + releaseTool.TrySetResult(); + } + } + + public sealed class FirstAnswer + { + public required int First { get; set; } + } + + public sealed class SecondAnswer + { + public required int Second { get; set; } + } + + public sealed class Inventory + { + public required int Count { get; set; } + public required string Color { get; set; } + } + + public sealed class CorrectionResult + { + public required int Answer { get; set; } + } + + public sealed class ToolAnswer + { + public required int Answer { get; set; } + public required string Contract { get; set; } + } + + [JsonSourceGenerationOptions(PropertyNamingPolicy = JsonKnownNamingPolicy.CamelCase)] + [JsonSerializable(typeof(Inventory))] + [JsonSerializable(typeof(FirstAnswer))] + [JsonSerializable(typeof(SecondAnswer))] + [JsonSerializable(typeof(CorrectionResult))] + [JsonSerializable(typeof(ToolAnswer))] + internal sealed partial class StructuredOutputE2EJsonContext : JsonSerializerContext; +} diff --git a/dotnet/test/Harness/ReplayProxy.cs b/dotnet/test/Harness/ReplayProxy.cs index 895ebccb87..1f906d72a5 100644 --- a/dotnet/test/Harness/ReplayProxy.cs +++ b/dotnet/test/Harness/ReplayProxy.cs @@ -250,7 +250,9 @@ public record ParsedHttpExchange( public record ChatCompletionRequest( string Model, List Messages, - List? Tools); + List? Tools, + [property: JsonPropertyName("tool_choice")] JsonElement? ToolChoice = null, + [property: JsonPropertyName("response_format")] JsonElement? ResponseFormat = null); public record ChatCompletionMessage( string Role, diff --git a/dotnet/test/Unit/ClientSessionLifetimeTests.cs b/dotnet/test/Unit/ClientSessionLifetimeTests.cs index 3e4973a192..e6764c0ef4 100644 --- a/dotnet/test/Unit/ClientSessionLifetimeTests.cs +++ b/dotnet/test/Unit/ClientSessionLifetimeTests.cs @@ -18,7 +18,7 @@ namespace GitHub.Copilot.Test.Unit; -public sealed class ClientSessionLifetimeTests +public sealed partial class ClientSessionLifetimeTests { private sealed record RpcRequestRecord(string Method, JsonElement Params); @@ -2376,6 +2376,11 @@ private sealed class FakeCopilotServer : IAsyncDisposable private bool _failRuntimeShutdown; private bool _failSessionCreate; private bool _failSessionSend; + private int _nextMessageId; + + public bool UniqueMessageIds { get; set; } + + public Func? BeforeSendResponse { get; set; } private FakeCopilotServer(TcpListener listener) { @@ -2476,7 +2481,7 @@ public async Task SendRequestAsync(string method, Dictionary data) + public Task SendSessionEventAsync(string sessionId, string type, Dictionary data, string? agentId = null) { var stream = _stream ?? throw new InvalidOperationException("Client is not connected."); var evt = new Dictionary @@ -2484,6 +2489,7 @@ public Task SendSessionEventAsync(string sessionId, string type, Dictionary new Dictionary @@ -2658,7 +2669,11 @@ private async Task HandleRequestAsync(Stream stream, JsonElement request, Cancel }, "session.send" => new Dictionary { - ["messageId"] = "message-1" + ["messageId"] = sendMessageId + }, + "session.sendMessages" => new Dictionary + { + ["messageIds"] = new[] { sendMessageId } }, "session.abort" => new Dictionary(), "session.getMessages" => new Dictionary diff --git a/dotnet/test/Unit/StructuredOutputTests.cs b/dotnet/test/Unit/StructuredOutputTests.cs new file mode 100644 index 0000000000..389166eeb1 --- /dev/null +++ b/dotnet/test/Unit/StructuredOutputTests.cs @@ -0,0 +1,542 @@ +/*--------------------------------------------------------------------------------------------- + * Copyright (c) Microsoft Corporation. All rights reserved. + *--------------------------------------------------------------------------------------------*/ + +#if NET8_0_OR_GREATER +using GitHub.Copilot.Rpc; +using System.Text.Json; +using System.Text.Json.Serialization; +using Xunit; + +namespace GitHub.Copilot.Test.Unit; + +public sealed partial class ClientSessionLifetimeTests +{ + [Theory] + [InlineData("session")] + [InlineData("rpc")] + [InlineData("batch")] + public async Task StructuredOutput_Raw_Format_Is_Forwarded_Without_Rewriting(string api) + { + await using var server = await FakeCopilotServer.StartAsync(); + await using var client = new CopilotClient(new CopilotClientOptions { Connection = RuntimeConnection.ForUri(server.Url) }); + await using var session = await client.CreateSessionAsync(new SessionConfig()); + using var document = JsonDocument.Parse("""{"type":"object","properties":{"value":{"type":"integer"}},"x-provider":{"anything":[true,42,null]}}"""); + var format = new ResponseFormatJsonSchema + { + JsonSchema = new JsonSchemaResponseFormat + { + Name = "answer", + Schema = document.RootElement.Clone(), + Strict = false, + Description = "An answer", + }, + }; + var options = new MessageOptions { Prompt = "Answer", ResponseSchema = document.RootElement.Clone() }; + Assert.Equal(options.ResponseSchema, options.Clone().ResponseSchema); + if (api == "batch") + { + await session.Rpc.SendMessagesAsync([new() { Prompt = "Answer" }], responseFormat: format); + } + else if (api == "rpc") + { + await session.Rpc.SendAsync("Answer", responseFormat: format); + } + else + { + await session.SendAsync(options); + } + var request = Assert.Single(server.Requests, r => r.Method == (api == "batch" ? "session.sendMessages" : "session.send")); + var wireFormat = request.Params.GetProperty("responseFormat"); + Assert.Equal("json_schema", wireFormat.GetProperty("type").GetString()); + var jsonSchema = wireFormat.GetProperty("jsonSchema"); + Assert.Equal(api == "session" ? "response" : "answer", jsonSchema.GetProperty("name").GetString()); + if (api == "session") + { + Assert.False(jsonSchema.TryGetProperty("description", out _)); + Assert.True(jsonSchema.GetProperty("strict").GetBoolean()); + } + else + { + Assert.Equal("An answer", jsonSchema.GetProperty("description").GetString()); + Assert.False(jsonSchema.GetProperty("strict").GetBoolean()); + } + Assert.Equal(document.RootElement.GetRawText(), jsonSchema.GetProperty("schema").GetRawText()); + + server.ClearRequests(); + await session.SendAsync("Ordinary text"); + Assert.False(Assert.Single(server.Requests, r => r.Method == "session.send").Params.TryGetProperty("responseFormat", out _)); + } + + [Fact] + public async Task StructuredOutput_Uses_Default_Custom_Tool_Serialization_Options() + { + await using var server = await FakeCopilotServer.StartAsync(); + await using var client = new CopilotClient(new CopilotClientOptions { Connection = RuntimeConnection.ForUri(server.Url) }); + await using var session = await client.CreateSessionAsync(new SessionConfig()); + var task = session.SendAndWaitAsync("Answer"); + if (!JsonSerializer.IsReflectionEnabledByDefault) + { + await Assert.ThrowsAsync(() => task); + Assert.DoesNotContain(server.Requests, request => request.Method == "session.send"); + return; + } + var request = await WaitForRequestAsync(server, "session.send"); + var properties = request.Params.GetProperty("responseFormat").GetProperty("jsonSchema").GetProperty("schema").GetProperty("properties"); + Assert.True(properties.TryGetProperty("answer_text", out _)); + Assert.True(properties.TryGetProperty("count", out _)); + await SendStructuredAnswerAsync(server, session, "message-1", """{"answer_text":"correct","count":42}"""); + var result = await task.WaitAsync(TimeSpan.FromSeconds(5)); + Assert.Equal("correct", result.Answer); + Assert.Equal(42, result.Count); + } + + [Theory] + [InlineData(false)] + [InlineData(true)] + public async Task StructuredOutput_Initializes_Fresh_Serializer_Options_Without_Mutating_Them(bool withResolver) + { + await using var server = await FakeCopilotServer.StartAsync(); + await using var client = new CopilotClient(new CopilotClientOptions { Connection = RuntimeConnection.ForUri(server.Url) }); + await using var session = await client.CreateSessionAsync(new SessionConfig()); + var options = new JsonSerializerOptions { PropertyNamingPolicy = JsonNamingPolicy.CamelCase }; + if (withResolver) + { + options.TypeInfoResolver = StructuredOutputJsonContext.Default; + } + var resolver = options.TypeInfoResolver; + var task = session.SendAndWaitAsync("Answer", options); + if (JsonSerializer.IsReflectionEnabledByDefault || withResolver) + { + var request = await WaitForRequestAsync(server, "session.send"); + var properties = request.Params.GetProperty("responseFormat").GetProperty("jsonSchema").GetProperty("schema").GetProperty("properties"); + Assert.True(properties.TryGetProperty("answer_text", out _)); + Assert.True(properties.TryGetProperty("count", out _)); + await SendStructuredAnswerAsync(server, session, "message-1", """{"answer_text":"correct","count":42}"""); + var result = await task.WaitAsync(TimeSpan.FromSeconds(5)); + Assert.Equal("correct", result.Answer); + Assert.Equal(42, result.Count); + } + else + { + await Assert.ThrowsAsync(() => task); + Assert.DoesNotContain(server.Requests, request => request.Method == "session.send"); + } + Assert.Same(resolver, options.TypeInfoResolver); + Assert.False(options.IsReadOnly); + } + + [Fact] + public async Task StructuredOutput_Infers_Schema_Using_Serialization_Contract() + { + await using var server = await FakeCopilotServer.StartAsync(); + await using var client = new CopilotClient(new CopilotClientOptions { Connection = RuntimeConnection.ForUri(server.Url) }); + await using var session = await client.CreateSessionAsync(new SessionConfig()); + var options = new MessageOptions { Prompt = "Answer", RequestHeaders = new Dictionary { ["x-test"] = "preserved" } }; + var task = session.SendAndWaitAsync(options, StructuredOutputJsonContext.Default.Options); + var request = await WaitForRequestAsync(server, "session.send"); + Assert.Null(options.ResponseSchema); + Assert.Equal("preserved", request.Params.GetProperty("requestHeaders").GetProperty("x-test").GetString()); + var format = request.Params.GetProperty("responseFormat").GetProperty("jsonSchema"); + Assert.True(format.GetProperty("strict").GetBoolean()); + var schema = format.GetProperty("schema"); + var properties = schema.GetProperty("properties"); + Assert.True(properties.TryGetProperty("answer_text", out _)); + Assert.True(properties.TryGetProperty("count", out _)); + Assert.True(properties.TryGetProperty("note", out var note)); + Assert.Contains("null", note.GetProperty("type").EnumerateArray().Select(t => t.GetString())); + Assert.False(schema.GetProperty("additionalProperties").GetBoolean()); + Assert.Equal(3, schema.GetProperty("required").GetArrayLength()); + await SendStructuredAnswerAsync(server, session, "message-1", """{"answer_text":"correct","count":42,"note":null}"""); + var result = await task.WaitAsync(TimeSpan.FromSeconds(5)); + Assert.Equal("correct", result.Answer); + Assert.Equal(42, result.Count); + Assert.Null(result.Note); + } + + [Theory] + [InlineData("not JSON")] + [InlineData("""{"answer_text":"wrong","count":"not a number"}""")] + [InlineData("null")] + [InlineData("""{"count":42}""")] + public async Task StructuredOutput_Rejects_Unparseable_Or_Null_Result(string content) + { + await using var server = await FakeCopilotServer.StartAsync(); + await using var client = new CopilotClient(new CopilotClientOptions { Connection = RuntimeConnection.ForUri(server.Url) }); + await using var session = await client.CreateSessionAsync(new SessionConfig()); + var task = session.SendAndWaitAsync("Answer", StructuredOutputJsonContext.Default.Options); + await WaitForRequestAsync(server, "session.send"); + await SendStructuredAnswerAsync(server, session, "message-1", content); + await Assert.ThrowsAsync(() => task.WaitAsync(TimeSpan.FromSeconds(5))); + } + + [Fact] + public async Task StructuredOutput_Rejects_Conflicting_Options_Before_Sending() + { + using var schema = JsonDocument.Parse("{}"); + await using var server = await FakeCopilotServer.StartAsync(); + await using var client = new CopilotClient(new CopilotClientOptions { Connection = RuntimeConnection.ForUri(server.Url) }); + await using var session = await client.CreateSessionAsync(new SessionConfig()); + await Assert.ThrowsAsync(() => session.SendAndWaitAsync( + new MessageOptions { Prompt = "Answer", Mode = "immediate" }, StructuredOutputJsonContext.Default.Options)); + await Assert.ThrowsAsync(() => session.SendAndWaitAsync( + new MessageOptions { Prompt = "Answer", ResponseSchema = schema.RootElement.Clone() }, StructuredOutputJsonContext.Default.Options)); + Assert.DoesNotContain(server.Requests, r => r.Method == "session.send"); + } + + [Theory] + [InlineData(false)] + [InlineData(true)] + public async Task StructuredOutput_Correlates_Concurrent_Queued_Sends(bool typed) + { + await using var server = await FakeCopilotServer.StartAsync(); + server.UniqueMessageIds = true; + await using var client = new CopilotClient(new CopilotClientOptions { Connection = RuntimeConnection.ForUri(server.Url) }); + await using var session = await client.CreateSessionAsync(new SessionConfig()); + async Task SendAsync(string prompt) + { + if (typed) + { + return await session.SendAndWaitAsync(prompt, StructuredOutputJsonContext.Default.Options); + } + using var schema = JsonDocument.Parse("""{"type":"object","properties":{"answer_text":{"type":"string"},"count":{"type":"integer"}},"required":["answer_text","count"],"additionalProperties":false}"""); + var message = await session.SendAndWaitAsync(new MessageOptions { Prompt = prompt, ResponseSchema = schema.RootElement.Clone() }); + Assert.NotNull(message); + return JsonSerializer.Deserialize(message.Data.Content, StructuredOutputJsonContext.Default.StructuredAnswer)!; + } + var first = SendAsync("First"); + await WaitForRequestAsync(server, "session.send"); + server.ClearRequests(); + var second = SendAsync("Second"); + await WaitForRequestAsync(server, "session.send"); + + await SendStructuredAnswerAsync(server, session, "message-1", """{"answer_text":"first","count":1}"""); + Assert.Equal("first", (await first.WaitAsync(TimeSpan.FromSeconds(5))).Answer); + Assert.False(second.IsCompleted); + await SendStructuredAnswerAsync(server, session, "message-2", """{"answer_text":"second","count":2}"""); + Assert.Equal("second", (await second.WaitAsync(TimeSpan.FromSeconds(5))).Answer); + } + + [Fact] + public async Task StructuredOutput_Buffers_Events_Before_Send_Response() + { + await using var server = await FakeCopilotServer.StartAsync(); + await using var client = new CopilotClient(new CopilotClientOptions { Connection = RuntimeConnection.ForUri(server.Url) }); + await using var session = await client.CreateSessionAsync(new SessionConfig()); + server.BeforeSendResponse = messageId => + SendStructuredAnswerAsync(server, session, messageId, """{"answer_text":"early","count":42}"""); + var result = await session.SendAndWaitAsync( + "Answer", StructuredOutputJsonContext.Default.Options, TimeSpan.FromSeconds(5)); + Assert.Equal("early", result.Answer); + } + + [Fact] + public async Task StructuredOutput_Ignores_Idle_Until_Own_Message_Is_Consumed() + { + await using var server = await FakeCopilotServer.StartAsync(); + await using var client = new CopilotClient(new CopilotClientOptions { Connection = RuntimeConnection.ForUri(server.Url) }); + await using var session = await client.CreateSessionAsync(new SessionConfig()); + var task = session.SendAndWaitAsync("Answer", StructuredOutputJsonContext.Default.Options); + await WaitForRequestAsync(server, "session.send"); + await server.SendSessionEventAsync(session.SessionId, "session.idle", new()); + await server.SendSessionEventAsync(session.SessionId, "session.error", new() + { + ["errorType"] = "provider", + ["message"] = "another turn failed", + }); + await SendStructuredAnswerAsync(server, session, "another-message", """{"answer_text":"wrong","count":0}"""); + await server.SendSessionEventAsync(session.SessionId, "user.message", new() + { + ["messageId"] = "message-1", + ["content"] = "Answer", + }, agentId: "subagent"); + await server.SendSessionEventAsync(session.SessionId, "session.idle", new()); + await SendStructuredAnswerAsync(server, session, "message-1", """{"answer_text":"correct","count":42}"""); + Assert.Equal("correct", (await task.WaitAsync(TimeSpan.FromSeconds(5))).Answer); + } + + [Fact] + public async Task StructuredOutput_Ignores_Subagent_Completion_And_Autopilot_Idle() + { + await using var server = await FakeCopilotServer.StartAsync(); + await using var client = new CopilotClient(new CopilotClientOptions { Connection = RuntimeConnection.ForUri(server.Url) }); + await using var session = await client.CreateSessionAsync(new SessionConfig()); + var task = session.SendAndWaitAsync("Answer", StructuredOutputJsonContext.Default.Options); + await WaitForRequestAsync(server, "session.send"); + await server.SendSessionEventAsync(session.SessionId, "user.message", new() + { + ["messageId"] = "message-1", + ["content"] = "Answer", + }); + await server.SendSessionEventAsync(session.SessionId, "session.idle", new(), agentId: "subagent"); + await server.SendSessionEventAsync(session.SessionId, "session.error", new() + { + ["errorType"] = "provider", + ["message"] = "subagent failed", + }, agentId: "subagent"); + await server.SendSessionEventAsync(session.SessionId, "session.idle", new() { ["mode"] = "autopilot" }); + await SendStructuredAnswerAsync(server, session, "message-1", """{"answer_text":"correct","count":42}"""); + Assert.Equal("correct", (await task.WaitAsync(TimeSpan.FromSeconds(5))).Answer); + } + + [Theory] + [InlineData(false)] + [InlineData(true)] + public async Task StructuredOutput_Rejects_Missing_Final_Response(bool toolOnly) + { + await using var server = await FakeCopilotServer.StartAsync(); + await using var client = new CopilotClient(new CopilotClientOptions { Connection = RuntimeConnection.ForUri(server.Url) }); + await using var session = await client.CreateSessionAsync(new SessionConfig()); + var task = session.SendAndWaitAsync("Answer", StructuredOutputJsonContext.Default.Options); + await WaitForRequestAsync(server, "session.send"); + await server.SendSessionEventAsync(session.SessionId, "user.message", new() + { + ["messageId"] = "message-1", + ["content"] = "Answer", + }); + if (toolOnly) + { + await server.SendSessionEventAsync(session.SessionId, "assistant.message", new() + { + ["messageId"] = "tool-message", + ["originatingMessageId"] = "message-1", + ["content"] = """{"answer_text":"not final","count":42}""", + ["toolRequests"] = new[] { new Dictionary { ["toolCallId"] = "tool-1", ["name"] = "terminal_tool", ["arguments"] = new Dictionary() } }, + }); + } + await server.SendSessionEventAsync(session.SessionId, "session.idle", new()); + var error = await Assert.ThrowsAsync(() => task.WaitAsync(TimeSpan.FromSeconds(5))); + Assert.Contains("without a final structured response", error.Message); + } + + [Theory] + [InlineData(false)] + [InlineData(true)] + public async Task StructuredOutput_Uses_Last_Correlated_Message_Not_Subagent_Or_Tool_Commentary(bool laterWorkAborted) + { + await using var server = await FakeCopilotServer.StartAsync(); + await using var client = new CopilotClient(new CopilotClientOptions { Connection = RuntimeConnection.ForUri(server.Url) }); + await using var session = await client.CreateSessionAsync(new SessionConfig()); + var task = session.SendAndWaitAsync("Answer", StructuredOutputJsonContext.Default.Options); + await WaitForRequestAsync(server, "session.send"); + await server.SendSessionEventAsync(session.SessionId, "user.message", new() + { + ["messageId"] = "message-1", + ["content"] = "Answer", + }); + foreach (var (origin, content) in new[] + { + ("message-1", "First I will inspect the inventory."), + ("message-1", """{"answer_text":"final","count":42}"""), + ("subagent-message", """{"answer_text":"wrong","count":0}"""), + ("unrelated-queued-message", """{"answer_text":"also wrong","count":0}"""), + }) + { + await server.SendSessionEventAsync(session.SessionId, "assistant.message", new() + { + ["messageId"] = Guid.NewGuid().ToString(), + ["originatingMessageId"] = origin, + ["content"] = content, + }); + } + await server.SendSessionEventAsync(session.SessionId, "assistant.message", new() + { + ["messageId"] = "subagent-output", + ["originatingMessageId"] = "message-1", + ["content"] = """{"answer_text":"subagent must not win","count":0}""", + }, agentId: "subagent-1"); + await server.SendSessionEventAsync(session.SessionId, "session.idle", new() { ["aborted"] = laterWorkAborted }); + if (laterWorkAborted) + { + var error = await Assert.ThrowsAsync(() => task.WaitAsync(TimeSpan.FromSeconds(5))); + Assert.Contains("aborted", error.Message); + } + else + { + Assert.Equal("final", (await task.WaitAsync(TimeSpan.FromSeconds(5))).Answer); + } + } + + [Theory] + [InlineData(false)] + [InlineData(true)] + public async Task StructuredOutput_Preserves_Timeout_And_Cancellation(bool cancel) + { + await using var server = await FakeCopilotServer.StartAsync(); + await using var client = new CopilotClient(new CopilotClientOptions { Connection = RuntimeConnection.ForUri(server.Url) }); + await using var session = await client.CreateSessionAsync(new SessionConfig()); + using var cts = new CancellationTokenSource(); + var task = session.SendAndWaitAsync("Answer", StructuredOutputJsonContext.Default.Options, + cancel ? TimeSpan.FromSeconds(10) : TimeSpan.FromMilliseconds(100), cts.Token); + await WaitForRequestAsync(server, "session.send"); + if (cancel) + { + cts.Cancel(); + await Assert.ThrowsAnyAsync(() => task); + } + else + { + await Assert.ThrowsAsync(() => task); + } + } + + [Theory] + [InlineData(false)] + [InlineData(true)] + public async Task StructuredOutput_Propagates_Rpc_And_Session_Errors(bool rpcError) + { + await using var server = await FakeCopilotServer.StartAsync(); + await using var client = new CopilotClient(new CopilotClientOptions { Connection = RuntimeConnection.ForUri(server.Url) }); + await using var session = await client.CreateSessionAsync(new SessionConfig()); + if (rpcError) + { + server.FailSessionSend(); + } + var task = session.SendAndWaitAsync("Answer", StructuredOutputJsonContext.Default.Options); + await WaitForRequestAsync(server, "session.send"); + if (!rpcError) + { + await server.SendSessionEventAsync(session.SessionId, "user.message", new() + { + ["messageId"] = "message-1", + ["content"] = "Answer", + }); + await server.SendSessionEventAsync(session.SessionId, "session.error", new() + { + ["errorType"] = "provider", + ["message"] = "structured output unsupported", + }); + var error = await Assert.ThrowsAsync(() => task.WaitAsync(TimeSpan.FromSeconds(5))); + Assert.Contains("structured output unsupported", error.Message); + } + else + { + var error = await Assert.ThrowsAsync(() => task.WaitAsync(TimeSpan.FromSeconds(5))); + Assert.Contains("session send failed", error.Message); + } + } + + [Fact] + public async Task StructuredOutput_Response_Does_Not_Hide_Later_Session_Errors() + { + await using var server = await FakeCopilotServer.StartAsync(); + await using var client = new CopilotClient(new CopilotClientOptions { Connection = RuntimeConnection.ForUri(server.Url) }); + await using var session = await client.CreateSessionAsync(new SessionConfig()); + var task = session.SendAndWaitAsync("Answer", StructuredOutputJsonContext.Default.Options); + await WaitForRequestAsync(server, "session.send"); + await server.SendSessionEventAsync(session.SessionId, "user.message", new() + { + ["messageId"] = "message-1", + ["content"] = "Answer", + }); + await server.SendSessionEventAsync(session.SessionId, "assistant.message", new() + { + ["messageId"] = "final-reply", + ["originatingMessageId"] = "message-1", + ["content"] = """{"answer_text":"correct","count":42}""", + }); + await server.SendSessionEventAsync(session.SessionId, "session.error", new() + { + ["errorType"] = "query", + ["message"] = "post-response failure", + }); + await server.SendSessionEventAsync(session.SessionId, "session.idle", new()); + var error = await Assert.ThrowsAsync(() => task.WaitAsync(TimeSpan.FromSeconds(5))); + Assert.Contains("post-response failure", error.Message); + } + + [Fact] + public async Task StructuredOutput_Can_Correlate_Without_User_Message_Event() + { + await using var server = await FakeCopilotServer.StartAsync(); + await using var client = new CopilotClient(new CopilotClientOptions { Connection = RuntimeConnection.ForUri(server.Url) }); + await using var session = await client.CreateSessionAsync(new SessionConfig()); + var task = session.SendAndWaitAsync("Answer", StructuredOutputJsonContext.Default.Options); + await WaitForRequestAsync(server, "session.send"); + await server.SendSessionEventAsync(session.SessionId, "assistant.message", new() + { + ["messageId"] = "assistant-result", + ["originatingMessageId"] = "message-1", + ["content"] = """{"answer_text":"correct","count":42}""", + }); + await server.SendSessionEventAsync(session.SessionId, "session.idle", new()); + Assert.Equal("correct", (await task.WaitAsync(TimeSpan.FromSeconds(5))).Answer); + } + + [Theory] + [InlineData(false)] + [InlineData(true)] + public async Task StructuredOutput_Rejects_When_Connection_Or_Session_Closes(bool disposeSession) + { + await using var server = await FakeCopilotServer.StartAsync(); + await using var client = new CopilotClient(new CopilotClientOptions { Connection = RuntimeConnection.ForUri(server.Url) }); + await using var session = await client.CreateSessionAsync(new SessionConfig()); + var task = session.SendAndWaitAsync("Answer", StructuredOutputJsonContext.Default.Options); + await WaitForRequestAsync(server, "session.send"); + if (disposeSession) + { + await session.DisposeAsync(); + } + else + { + server.CloseConnection(); + } + try + { + await Assert.ThrowsAsync(() => task.WaitAsync(TimeSpan.FromSeconds(5))); + } + finally + { + // Graceful cleanup cannot wait for a peer whose transport was deliberately closed. + await client.ForceStopAsync(); + } + } + + [Fact] + public async Task StructuredOutput_Timeout_Includes_Send_Acknowledgement() + { + await using var server = await FakeCopilotServer.StartAsync(); + await using var client = new CopilotClient(new CopilotClientOptions { Connection = RuntimeConnection.ForUri(server.Url) }); + await using var session = await client.CreateSessionAsync(new SessionConfig()); + var release = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + server.BeforeSendResponse = _ => release.Task; + try + { + await Assert.ThrowsAsync(() => session.SendAndWaitAsync( + "Answer", StructuredOutputJsonContext.Default.Options, TimeSpan.FromMilliseconds(100))); + } + finally + { + release.TrySetResult(); + } + } + + private static async Task SendStructuredAnswerAsync(FakeCopilotServer server, CopilotSession session, string messageId, string content) + { + await server.SendSessionEventAsync(session.SessionId, "user.message", new() + { + ["messageId"] = messageId, + ["content"] = "Answer", + }); + await server.SendSessionEventAsync(session.SessionId, "assistant.message", new() + { + ["messageId"] = Guid.NewGuid().ToString(), + ["originatingMessageId"] = messageId, + ["content"] = content, + }); + await server.SendSessionEventAsync(session.SessionId, "session.idle", new()); + } + + public sealed class StructuredAnswer + { + [JsonPropertyName("answer_text")] + public required string Answer { get; set; } + public int Count { get; set; } + public string? Note { get; set; } + } + + [JsonSourceGenerationOptions(PropertyNamingPolicy = JsonKnownNamingPolicy.CamelCase)] + [JsonSerializable(typeof(StructuredAnswer))] + internal sealed partial class StructuredOutputJsonContext : JsonSerializerContext; +} +#endif diff --git a/go/README.md b/go/README.md index cbbb84b78c..fd8bdd245a 100644 --- a/go/README.md +++ b/go/README.md @@ -529,6 +529,38 @@ lookupIssue := copilot.DefineTool("lookup_issue", "Fetch issue details", lookupIssue.Defer = copilot.ToolDeferAuto ``` +## Structured output (experimental) + +Use the package-level generic helper (Go does not support generic methods): + +```go +type Inventory struct { + Count int `json:"count"` + Color string `json:"color"` +} + +inventory, err := copilot.SendAndWait[Inventory](ctx, session, copilot.MessageOptions{ + Prompt: "Call get_inventory, then report the widget count and color.", +}) +``` + +This derives the schema using the same `jsonschema-go` generator as `DefineTool` +and unmarshals the final JSON into `Inventory`. Unmarshaling is not full JSON +Schema validation. For an explicit schema, set `MessageOptions.ResponseSchema` +on `session.Send` or `session.SendAndWait`; the latter returns the message event. +The generic helper rejects an explicit schema or immediate delivery. + +The schema lasts for one run, including tools, steering, and stop-hook corrections. +Independent sends and subagents do not inherit it; streaming stays text. +Structured waits select the last correlated root assistant message without tool +requests at non-autopilot idle. Concurrent waits retain their own results, though +queued work can delay idle. Aborted runs, session errors after the run starts, +and missing final output fail. Context cancellation stops waiting, not agent work. + +Provider schema restrictions apply, and supplied schemas are forwarded unchanged. +Low-level `session.RPC.Send` and `session.RPC.SendMessages` expose the full +`rpc.ResponseFormat` options, including name, description, and strictness. + ## Streaming Enable streaming to receive assistant response chunks as they're generated: diff --git a/go/internal/e2e/structured_output_e2e_test.go b/go/internal/e2e/structured_output_e2e_test.go new file mode 100644 index 0000000000..895e8bd2a7 --- /dev/null +++ b/go/internal/e2e/structured_output_e2e_test.go @@ -0,0 +1,548 @@ +package e2e + +import ( + "context" + "encoding/json" + "strings" + "sync" + "sync/atomic" + "testing" + "time" + + copilot "github.com/github/copilot-sdk/go" + "github.com/github/copilot-sdk/go/internal/e2e/testharness" + "github.com/github/copilot-sdk/go/rpc" +) + +type outputInventory struct { + Count int `json:"count"` + Color string `json:"color"` +} + +type outputAnswer struct { + Answer int `json:"answer"` +} + +func structuredSessionConfig(proxy string) *copilot.SessionConfig { + return &copilot.SessionConfig{ + Model: "gpt-4.1", AvailableTools: []string{}, + OnPermissionRequest: copilot.PermissionHandler.ApproveAll, + Provider: &copilot.ProviderConfig{ + Type: "openai", WireAPI: "completions", BaseURL: proxy, + ModelID: "gpt-4.1", WireModel: "gpt-4.1", APIKey: "fake-token-for-e2e-tests", + Headers: map[string]string{ + "Copilot-Integration-Id": "copilot-developer-cli", + "Copilot-Harness-Id": "copilot-sdk", + "X-GitHub-Api-Version": "2026-08-01", + }, + }, + } +} + +func TestStructuredOutputE2E(t *testing.T) { + harness := testharness.NewTestContext(t) + client := harness.NewClient() + t.Cleanup(func() { client.ForceStop() }) + + t.Run("infers_typed_result_after_custom_tool", func(t *testing.T) { + harness.ConfigureForTest(t) + ctx, cancel := context.WithTimeout(t.Context(), 45*time.Second) + defer cancel() + var calls atomic.Int32 + config := structuredSessionConfig(harness.ProxyURL) + streaming := true + config.Streaming = &streaming + config.Tools = []copilot.Tool{copilot.DefineTool("get_inventory", "Get the current widget inventory.", + func(_ struct{}, _ copilot.ToolInvocation) (string, error) { + calls.Add(1) + return "The inventory contains 42 red widgets.", nil + })} + session, err := client.CreateSession(ctx, config) + if err != nil { + t.Fatal(err) + } + var deltas atomic.Int32 + unsubscribe := session.On(func(event copilot.SessionEvent) { + if _, ok := event.Data.(*copilot.AssistantMessageDeltaData); ok { + deltas.Add(1) + } + }) + defer unsubscribe() + result, err := copilot.SendAndWait[outputInventory](ctx, session, copilot.MessageOptions{ + Prompt: "Call get_inventory, then report the widget count and color.", + }) + if err != nil || result != (outputInventory{42, "red"}) || calls.Load() == 0 { + t.Fatalf("unexpected result %+v, calls %d, error %v", result, calls.Load(), err) + } + if deltas.Load() == 0 { + t.Fatal("typed wait did not stream text updates") + } + ordinary, err := session.SendAndWait(ctx, copilot.MessageOptions{Prompt: "Now reply with exactly the plain text HELLO, not JSON."}) + if err != nil { + t.Fatal(err) + } + if ordinary == nil || strings.TrimSpace(ordinary.Data.(*copilot.AssistantMessageData).Content) != "HELLO" { + t.Fatalf("schema leaked into ordinary follow-up: %v", ordinary) + } + exchanges, err := harness.GetExchanges() + if err != nil || len(exchanges) < 3 { + t.Fatalf("expected tool, final, and ordinary requests: %d, error %v", len(exchanges), err) + } + for _, exchange := range exchanges[:len(exchanges)-1] { + format := exchange.Request.ResponseFormat + contract, ok := format["json_schema"].(map[string]any) + if format["type"] != "json_schema" || !ok || contract["strict"] != true { + t.Fatalf("missing provider-native schema: %+v", format) + } + schema := contract["schema"].(map[string]any) + properties := schema["properties"].(map[string]any) + if schema["additionalProperties"] != false || + properties["count"].(map[string]any)["type"] != "integer" || + properties["color"].(map[string]any)["type"] != "string" { + t.Fatalf("incorrect inferred inventory schema: %+v", schema) + } + } + if exchanges[len(exchanges)-1].Request.ResponseFormat != nil { + t.Fatal("response_format leaked into ordinary follow-up") + } + }) + + t.Run("typed_wait_returns_stop_hook_correction", func(t *testing.T) { + harness.ConfigureForTest(t) + ctx, cancel := context.WithTimeout(t.Context(), 45*time.Second) + defer cancel() + var stops atomic.Int32 + config := structuredSessionConfig(harness.ProxyURL) + config.Hooks = &copilot.SessionHooks{OnAgentStop: func(_ copilot.AgentStopHookInput, _ copilot.HookInvocation) (*copilot.AgentStopHookOutput, error) { + if stops.Add(1) == 1 { + return &copilot.AgentStopHookOutput{Decision: "block", Reason: "Correct the answer to 99, not 42. Do not use tools."}, nil + } + return nil, nil + }} + session, err := client.CreateSession(ctx, config) + if err != nil { + t.Fatal(err) + } + var mu sync.Mutex + var replies []*copilot.AssistantMessageData + unsubscribe := session.On(func(event copilot.SessionEvent) { + if reply, ok := event.Data.(*copilot.AssistantMessageData); ok && (event.AgentID == nil || *event.AgentID == "") { + mu.Lock() + replies = append(replies, reply) + mu.Unlock() + } + }) + defer unsubscribe() + result, err := copilot.SendAndWait[outputAnswer](ctx, session, copilot.MessageOptions{Prompt: "What is 19 + 23? Do not use tools."}) + if err != nil || result.Answer != 99 || stops.Load() != 2 { + t.Fatalf("result %+v, stops %d, error %v", result, stops.Load(), err) + } + mu.Lock() + defer mu.Unlock() + if len(replies) != 2 || replies[0].OriginatingMessageID == nil || *replies[0].OriginatingMessageID == "" || + replies[1].OriginatingMessageID == nil || *replies[0].OriginatingMessageID != *replies[1].OriginatingMessageID { + t.Fatalf("expected two replies with the same nonempty origin: %+v", replies) + } + for i, expected := range []int{42, 99} { + var answer outputAnswer + if err := json.Unmarshal([]byte(replies[i].Content), &answer); err != nil || answer.Answer != expected { + t.Fatalf("reply %d: %+v, error %v", i, answer, err) + } + } + }) + + t.Run("send_selects_correlated_response_after_idle", func(t *testing.T) { + harness.ConfigureForTest(t) + ctx, cancel := context.WithTimeout(t.Context(), 45*time.Second) + defer cancel() + entered, release := make(chan struct{}), make(chan struct{}) + var releaseOnce sync.Once + defer releaseOnce.Do(func() { close(release) }) + config := structuredSessionConfig(harness.ProxyURL) + config.Tools = []copilot.Tool{copilot.DefineTool("read_inventory", "Read the current widget count and color.", + func(_ struct{}, _ copilot.ToolInvocation) (string, error) { + return "The inventory contains 42 red widgets.", nil + })} + config.Hooks = &copilot.SessionHooks{OnAgentStop: func(_ copilot.AgentStopHookInput, _ copilot.HookInvocation) (*copilot.AgentStopHookOutput, error) { + close(entered) + select { + case <-release: + return nil, nil + case <-ctx.Done(): + return nil, ctx.Err() + } + }} + session, err := client.CreateSession(ctx, config) + if err != nil { + t.Fatal(err) + } + idle := make(chan struct{}, 1) + failures := make(chan string, 1) + var mu sync.Mutex + var replies []*copilot.AssistantMessageData + unsubscribe := session.On(func(event copilot.SessionEvent) { + if event.AgentID != nil && *event.AgentID != "" { + return + } + switch data := event.Data.(type) { + case *copilot.AssistantMessageData: + mu.Lock() + replies = append(replies, data) + mu.Unlock() + case *copilot.SessionIdleData: + select { + case idle <- struct{}{}: + default: + } + case *copilot.SessionErrorData: + select { + case failures <- data.Message: + default: + } + } + }) + defer unsubscribe() + schema := map[string]any{ + "type": "object", "properties": map[string]any{"count": map[string]any{"type": "integer"}, "color": map[string]any{"type": "string"}}, + "required": []string{"count", "color"}, "additionalProperties": false, + } + origin, err := session.Send(ctx, copilot.MessageOptions{ + Prompt: "Call read_inventory once, then report the current widget count and color.", ResponseSchema: schema, + }) + if err != nil { + t.Fatal(err) + } + select { + case <-entered: + case err := <-failures: + t.Fatal(err) + case <-ctx.Done(): + t.Fatal(ctx.Err()) + } + select { + case <-idle: + t.Fatal("idle before stop hook finished") + default: + } + releaseOnce.Do(func() { close(release) }) + select { + case <-idle: + case err := <-failures: + t.Fatal(err) + case <-ctx.Done(): + t.Fatal(ctx.Err()) + } + mu.Lock() + defer mu.Unlock() + if len(replies) < 2 { + t.Fatalf("expected tool and final messages, got %d", len(replies)) + } + final := replies[len(replies)-1] + var inventory outputInventory + if final.OriginatingMessageID == nil || *final.OriginatingMessageID != origin || len(final.ToolRequests) != 0 { + t.Fatalf("invalid final message: %+v", final) + } + if err := json.Unmarshal([]byte(final.Content), &inventory); err != nil || inventory != (outputInventory{42, "red"}) { + t.Fatalf("result %+v, error %v", inventory, err) + } + if len(replies[0].ToolRequests) == 0 { + t.Fatal("missing intermediate tool request") + } + }) + + t.Run("rejects_invalid_formats_before_admission", func(t *testing.T) { + harness.ConfigureForTest(t) + ctx, cancel := context.WithTimeout(t.Context(), 45*time.Second) + defer cancel() + session, err := client.CreateSession(ctx, structuredSessionConfig(harness.ProxyURL)) + if err != nil { + t.Fatal(err) + } + var admitted atomic.Bool + unsubscribe := session.On(func(event copilot.SessionEvent) { + switch event.Data.(type) { + case *copilot.UserMessageData, *copilot.SessionErrorData: + admitted.Store(true) + } + }) + defer unsubscribe() + if _, err := copilot.SendAndWait[outputAnswer](ctx, session, copilot.MessageOptions{Prompt: "Must not be admitted", Mode: "immediate"}); err == nil { + t.Fatal("typed immediate steering should be rejected") + } + schema := map[string]any{"type": "object", "description": strings.Repeat("x", 32*1024*1024)} + if _, err := session.SendAndWait(ctx, copilot.MessageOptions{Prompt: "Must not be admitted", ResponseSchema: schema}); err == nil || !strings.Contains(err.Error(), "32 MiB") { + t.Fatalf("expected schema size rejection: %v", err) + } + if _, err := session.RPC.SendMessages(ctx, &rpc.SendMessagesRequest{ + Messages: []rpc.SendMessageItem{}, + ResponseFormat: &rpc.ResponseFormat{Type: rpc.ResponseFormatTypeJSONSchema, JSONSchema: rpc.JSONSchemaResponseFormat{Name: "response", Schema: schema}}, + }); err == nil || !strings.Contains(err.Error(), "32 MiB") { + t.Fatalf("expected batch schema size rejection: %v", err) + } + pending, err := session.RPC.Queue.PendingItems(ctx) + if err != nil || len(pending.Items) != 0 || admitted.Load() { + t.Fatalf("rejected message was admitted: queue %+v, event %v, error %v", pending, admitted.Load(), err) + } + exchanges, err := harness.GetExchanges() + if err != nil || len(exchanges) != 0 { + t.Fatalf("unexpected provider calls: %d, error %v", len(exchanges), err) + } + }) + + t.Run("sends_explicit_schema_for_message_and_batch", func(t *testing.T) { + harness.ConfigureForTest(t) + ctx, cancel := context.WithTimeout(t.Context(), 45*time.Second) + defer cancel() + session, err := client.CreateSession(ctx, structuredSessionConfig(harness.ProxyURL)) + if err != nil { + t.Fatal(err) + } + var schema map[string]any + if err := json.Unmarshal([]byte(`{"type":"object","properties":{"count":{"type":"integer"},"color":{"type":"string"}},"required":["count","color"],"additionalProperties":false}`), &schema); err != nil { + t.Fatal(err) + } + completed := make(chan struct{}, 1) + var mu sync.Mutex + var replies []*copilot.AssistantMessageData + unsubscribe := session.On(func(event copilot.SessionEvent) { + if event.AgentID != nil && *event.AgentID != "" { + return + } + switch data := event.Data.(type) { + case *copilot.AssistantMessageData: + mu.Lock() + replies = append(replies, data) + mu.Unlock() + case *copilot.SessionIdleData: + select { + case completed <- struct{}{}: + default: + } + } + }) + defer unsubscribe() + strict := true + accepted, err := session.RPC.SendMessages(ctx, &rpc.SendMessagesRequest{ + Messages: []rpc.SendMessageItem{ + {Prompt: "There are 42 red widgets in stock."}, + {Prompt: "Report the widget count and color."}, + }, + ResponseFormat: &rpc.ResponseFormat{Type: rpc.ResponseFormatTypeJSONSchema, + JSONSchema: rpc.JSONSchemaResponseFormat{Name: "inventory", Schema: schema, Strict: &strict}}, + }) + if err != nil { + t.Fatal(err) + } + select { + case <-completed: + case <-ctx.Done(): + t.Fatal(ctx.Err()) + } + mu.Lock() + var content string + for _, reply := range replies { + if reply.OriginatingMessageID != nil && *reply.OriginatingMessageID == accepted.MessageIDs[len(accepted.MessageIDs)-1] { + content = reply.Content + } + } + mu.Unlock() + var inventory outputInventory + if err := json.Unmarshal([]byte(content), &inventory); err != nil { + t.Fatal(err) + } + if inventory != (outputInventory{42, "red"}) { + t.Fatalf("unexpected batch result: %+v", inventory) + } + raw, err := session.SendAndWait(ctx, copilot.MessageOptions{ + Prompt: "The inventory now has 21 blue widgets. Report the new count and color.", ResponseSchema: schema, + }) + if err != nil { + t.Fatal(err) + } + if err := json.Unmarshal([]byte(raw.Data.(*copilot.AssistantMessageData).Content), &inventory); err != nil { + t.Fatal(err) + } + if inventory != (outputInventory{21, "blue"}) { + t.Fatalf("unexpected raw result: %+v", inventory) + } + }) + + t.Run("typed_wait_returns_stop_hook_correction_after_terminal_tool", func(t *testing.T) { + harness.ConfigureForTest(t) + ctx, cancel := context.WithTimeout(t.Context(), 45*time.Second) + defer cancel() + var calls, stops atomic.Int32 + tool := copilot.DefineTool("lookup_number", "Return the number needed for the calculation.", + func(_ struct{}, _ copilot.ToolInvocation) (int, error) { calls.Add(1); return 58, nil }) + tool.IsTerminal, tool.SkipPermission = true, true + config := structuredSessionConfig(harness.ProxyURL) + config.Tools = []copilot.Tool{tool} + config.Hooks = &copilot.SessionHooks{OnAgentStop: func(_ copilot.AgentStopHookInput, _ copilot.HookInvocation) (*copilot.AgentStopHookOutput, error) { + if stops.Add(1) == 1 { + return &copilot.AgentStopHookOutput{Decision: "block", Reason: "Correct the answer to 99, not 63. Do not use tools."}, nil + } + return nil, nil + }} + session, err := client.CreateSession(ctx, config) + if err != nil { + t.Fatal(err) + } + result, err := copilot.SendAndWait[outputAnswer](ctx, session, copilot.MessageOptions{ + Prompt: "Call lookup_number exactly once, then add 5 to the returned number. Do not guess its result.", + }) + if err != nil || result.Answer != 99 || calls.Load() != 1 || stops.Load() != 2 { + t.Fatalf("result %+v, calls %d, stops %d, error %v", result, calls.Load(), stops.Load(), err) + } + }) + + t.Run("typed_wait_returns_late_steering_response", func(t *testing.T) { + harness.ConfigureForTest(t) + ctx, cancel := context.WithTimeout(t.Context(), 45*time.Second) + defer cancel() + var session *copilot.Session + var stops atomic.Int32 + config := structuredSessionConfig(harness.ProxyURL) + config.Hooks = &copilot.SessionHooks{OnAgentStop: func(_ copilot.AgentStopHookInput, _ copilot.HookInvocation) (*copilot.AgentStopHookOutput, error) { + if stops.Add(1) == 1 { + _, err := session.Send(ctx, copilot.MessageOptions{Prompt: "Change the answer to 99. Do not use tools.", Mode: "immediate"}) + return nil, err + } + return nil, nil + }} + var err error + session, err = client.CreateSession(ctx, config) + if err != nil { + t.Fatal(err) + } + result, err := copilot.SendAndWait[outputAnswer](ctx, session, copilot.MessageOptions{Prompt: "What is 19 + 23? Do not use tools."}) + if err != nil || result.Answer != 99 || stops.Load() != 2 { + t.Fatalf("result %+v, stops %d, error %v", result, stops.Load(), err) + } + }) + + t.Run("typed_result_after_terminal_tool_and_steering", func(t *testing.T) { + harness.ConfigureForTest(t) + ctx, cancel := context.WithTimeout(t.Context(), 45*time.Second) + defer cancel() + var session *copilot.Session + var calls atomic.Int32 + tool := copilot.DefineTool("lookup_number", "Return the number needed for the calculation.", + func(_ struct{}, _ copilot.ToolInvocation) (int, error) { + calls.Add(1) + _, err := session.Send(ctx, copilot.MessageOptions{ + Prompt: "Continue with the original calculation. Do not call any more tools.", Mode: "immediate", + }) + return 58, err + }) + tool.IsTerminal, tool.SkipPermission = true, true + config := structuredSessionConfig(harness.ProxyURL) + config.Tools = []copilot.Tool{tool} + var err error + session, err = client.CreateSession(ctx, config) + if err != nil { + t.Fatal(err) + } + type toolAnswer struct { + Answer int `json:"answer"` + Contract string `json:"contract"` + } + result, err := copilot.SendAndWait[toolAnswer](ctx, session, copilot.MessageOptions{ + Prompt: "Call lookup_number exactly once, then add 5 to the returned number. Do not guess its result.", + }) + if err != nil || result != (toolAnswer{63, "typed_tool"}) || calls.Load() != 1 { + t.Fatalf("result %+v, calls %d, error %v", result, calls.Load(), err) + } + exchanges, err := harness.GetExchanges() + if err != nil || len(exchanges) < 2 { + t.Fatalf("expected tool and final requests: %d, error %v", len(exchanges), err) + } + for _, exchange := range exchanges[1:] { + if string(exchange.Request.ToolChoice) != `"none"` { + t.Fatalf("terminal tool did not disable tools: %s", exchange.Request.ToolChoice) + } + } + for _, exchange := range exchanges { + if exchange.Request.ResponseFormat["type"] != "json_schema" { + t.Fatal("steering lost the active output schema") + } + } + }) + + t.Run("concurrent_typed_sends_return_their_own_results", func(t *testing.T) { + harness.ConfigureForTest(t) + ctx, cancel := context.WithTimeout(t.Context(), 45*time.Second) + defer cancel() + entered, release := make(chan struct{}), make(chan struct{}) + var releaseOnce sync.Once + defer releaseOnce.Do(func() { close(release) }) + config := structuredSessionConfig(harness.ProxyURL) + config.Tools = []copilot.Tool{copilot.DefineTool("first_number", "Get the number for the first question.", + func(_ struct{}, _ copilot.ToolInvocation) (int, error) { + close(entered) + select { + case <-release: + return 42, nil + case <-ctx.Done(): + return 0, ctx.Err() + } + })} + session, err := client.CreateSession(ctx, config) + if err != nil { + t.Fatal(err) + } + type firstAnswer struct { + First int `json:"first"` + } + type secondAnswer struct { + Second int `json:"second"` + } + type outcome struct { + value int + err error + } + first, second := make(chan outcome, 1), make(chan outcome, 1) + go func() { + result, err := copilot.SendAndWait[firstAnswer](ctx, session, copilot.MessageOptions{Prompt: "Call first_number exactly once and report its returned number."}) + first <- outcome{result.First, err} + }() + select { + case <-entered: + case result := <-first: + t.Fatalf("tool was not entered: %+v", result) + case <-ctx.Done(): + t.Fatal(ctx.Err()) + } + go func() { + result, err := copilot.SendAndWait[secondAnswer](ctx, session, copilot.MessageOptions{Prompt: "What is 30 + 7? Do not use tools."}) + second <- outcome{result.Second, err} + }() + for { + pending, err := session.RPC.Queue.PendingItems(ctx) + if err != nil { + t.Fatal(err) + } + if len(pending.Items) > 0 { + break + } + select { + case <-time.After(10 * time.Millisecond): + case <-ctx.Done(): + t.Fatal(ctx.Err()) + } + } + releaseOnce.Do(func() { close(release) }) + for _, expected := range []struct { + result <-chan outcome + value int + }{{first, 42}, {second, 37}} { + select { + case result := <-expected.result: + if result.err != nil || result.value != expected.value { + t.Fatalf("unexpected result: %+v", result) + } + case <-ctx.Done(): + t.Fatal(ctx.Err()) + } + } + }) +} diff --git a/go/internal/e2e/testharness/proxy.go b/go/internal/e2e/testharness/proxy.go index f120a0f8ba..f3dc1c4c09 100644 --- a/go/internal/e2e/testharness/proxy.go +++ b/go/internal/e2e/testharness/proxy.go @@ -283,9 +283,11 @@ type ParsedHttpExchange struct { // ChatCompletionRequest represents an OpenAI chat completion request. type ChatCompletionRequest struct { - Model string `json:"model"` - Messages []ChatCompletionMessage `json:"messages"` - Tools []ChatCompletionTool `json:"tools,omitempty"` + ToolChoice json.RawMessage `json:"tool_choice,omitempty"` + ResponseFormat map[string]any `json:"response_format,omitempty"` + Model string `json:"model"` + Messages []ChatCompletionMessage `json:"messages"` + Tools []ChatCompletionTool `json:"tools,omitempty"` } // ChatCompletionMessage represents a message in the chat completion request. diff --git a/go/internal/ffihost/ffihost.go b/go/internal/ffihost/ffihost.go index c98b5f4847..5d5161e833 100644 --- a/go/internal/ffihost/ffihost.go +++ b/go/internal/ffihost/ffihost.go @@ -208,9 +208,7 @@ type Host struct { } func (h *Host) rearmForeignSignalHandlers() { - if h.cliEntrypoint != "" { - rearmForeignSignalHandlers(h.lib.handle) - } + rearmForeignSignalHandlers(h.lib.handle) } // PrepareForChildProcessWait repairs signal handlers that the embedded runtime diff --git a/go/internal/ffihost/sigonstack_linux_test.go b/go/internal/ffihost/sigonstack_linux_test.go index 2f40026899..393d590d0c 100644 --- a/go/internal/ffihost/sigonstack_linux_test.go +++ b/go/internal/ffihost/sigonstack_linux_test.go @@ -39,6 +39,15 @@ func TestRearmForeignSignalHandlersAddsOnStack(t *testing.T) { } func TestHostRearmsSignalHandlersAroundNativeOperations(t *testing.T) { + for _, entrypoint := range []string{"", "copilot"} { + t.Run("entrypoint="+entrypoint, func(t *testing.T) { + testHostRearmsSignalHandlers(t, entrypoint) + }) + } +} + +func testHostRearmsSignalHandlers(t *testing.T, entrypoint string) { + t.Helper() signals := make(chan os.Signal, 1) signal.Notify(signals, syscall.SIGUSR1) defer signal.Stop(signals) @@ -50,7 +59,7 @@ func TestHostRearmsSignalHandlersAroundNativeOperations(t *testing.T) { defer linuxSetSigaction(int(syscall.SIGUSR1), &original) host := &Host{ - cliEntrypoint: "copilot", + cliEntrypoint: entrypoint, lib: &ffiLibrary{ hostStart: func(unsafe.Pointer, uintptr, unsafe.Pointer, uintptr) uint32 { return 1 diff --git a/go/message_source_test.go b/go/message_source_test.go index 0ff1d42545..e813617b9e 100644 --- a/go/message_source_test.go +++ b/go/message_source_test.go @@ -198,6 +198,10 @@ func TestSession_SendAndWaitMessageSource(t *testing.T) { } func captureMessageSourceRequest(t *testing.T, rpcError *jsonrpc2.Error, events []SessionEvent, invoke func(*Session)) map[string]any { + return captureSessionSendRequest(t, rpcError, events, false, invoke) +} + +func captureSessionSendRequest(t *testing.T, rpcError *jsonrpc2.Error, events []SessionEvent, beforeResponse bool, invoke func(*Session)) map[string]any { t.Helper() stdinR, stdinW := io.Pipe() @@ -246,12 +250,35 @@ func captureMessageSourceRequest(t *testing.T, rpcError *jsonrpc2.Error, events errCh <- err return } + if beforeResponse && len(events) > 0 { + processed := make(chan struct{}) + count := 0 + unsubscribe := session.On(func(SessionEvent) { + count++ + if count == len(events) { + close(processed) + } + }) + for _, event := range events { + session.dispatchEvent(event) + } + select { + case <-processed: + case <-time.After(time.Second): + errCh <- fmt.Errorf("pre-admission events were not dispatched") + unsubscribe() + return + } + unsubscribe() + } if _, err := fmt.Fprintf(stdoutW, "Content-Length: %d\r\n\r\n%s", len(data), data); err != nil { errCh <- err return } - for _, event := range events { - session.dispatchEvent(event) + if !beforeResponse { + for _, event := range events { + session.dispatchEvent(event) + } } paramsCh <- request.Params }() diff --git a/go/session.go b/go/session.go index 2dd2a8a9a1..4b279f3f98 100644 --- a/go/session.go +++ b/go/session.go @@ -441,6 +441,15 @@ func (s *Session) Send(ctx context.Context, options MessageOptions) (string, err Tracestate: tracestate, RequestHeaders: options.RequestHeaders, } + if options.ResponseSchema != nil { + strict := true + req.ResponseFormat = &rpc.ResponseFormat{ + Type: "json_schema", + JSONSchema: rpc.JSONSchemaResponseFormat{ + Name: "response", Schema: options.ResponseSchema, Strict: &strict, + }, + } + } result, err := s.client.Request(ctx, "session.send", req) if err != nil { @@ -492,6 +501,9 @@ func (s *Session) SendPrompt(ctx context.Context, prompt string) (string, error) // } // } func (s *Session) SendAndWait(ctx context.Context, options MessageOptions) (*SessionEvent, error) { + if options.ResponseSchema != nil { + return s.sendAndWaitStructured(ctx, options) + } if _, ok := ctx.Deadline(); !ok { var cancel context.CancelFunc ctx, cancel = context.WithTimeout(ctx, 60*time.Second) diff --git a/go/structured_output.go b/go/structured_output.go new file mode 100644 index 0000000000..73cecb3f26 --- /dev/null +++ b/go/structured_output.go @@ -0,0 +1,159 @@ +package copilot + +import ( + "bytes" + "context" + "encoding/json" + "fmt" + "reflect" + "strings" + "sync" + "time" + + "github.com/google/jsonschema-go/jsonschema" +) + +// SendAndWait infers a JSON Schema for T using the same generator as DefineTool, +// sends a message, and unmarshals its final correlated root response at session idle. +// Go does not support generic methods, so this is a package-level function. +// Options must not specify ResponseSchema or immediate delivery. Streaming events +// remain text. Cancellation stops waiting, not the agent; errors are session-scoped. +// Provider schema restrictions apply. Unmarshaling is not full JSON Schema validation. +func SendAndWait[T any](ctx context.Context, session *Session, options MessageOptions) (T, error) { + var result T + if session == nil { + return result, fmt.Errorf("session must not be nil") + } + if options.ResponseSchema != nil || options.Mode == "immediate" { + return result, fmt.Errorf("typed structured output cannot specify ResponseSchema or immediate delivery") + } + schema, err := jsonschema.ForType(reflect.TypeFor[T](), nil) + if err != nil { + return result, fmt.Errorf("infer response schema: %w", err) + } + encoded, err := json.Marshal(schema) + if err != nil { + return result, fmt.Errorf("marshal response schema: %w", err) + } + if err := json.Unmarshal(encoded, &options.ResponseSchema); err != nil { + return result, fmt.Errorf("decode response schema: %w", err) + } + event, err := session.SendAndWait(ctx, options) + if err != nil { + return result, err + } + content := event.Data.(*AssistantMessageData).Content + if bytes.Equal(bytes.TrimSpace([]byte(content)), []byte("null")) { + return result, fmt.Errorf("structured response was JSON null, not a result") + } + if err := json.Unmarshal([]byte(content), &result); err != nil { + return result, fmt.Errorf("decode structured response: %w", err) + } + return result, nil +} + +func (s *Session) sendAndWaitStructured(ctx context.Context, options MessageOptions) (*SessionEvent, error) { + if _, ok := ctx.Deadline(); !ok { + var cancel context.CancelFunc + ctx, cancel = context.WithTimeout(ctx, 60*time.Second) + defer cancel() + } + ctx, cancelWait := context.WithCancelCause(ctx) + defer cancelWait(nil) + go func() { + select { + case <-s.eventDone: + cancelWait(fmt.Errorf("session closed before structured output completed")) + case <-ctx.Done(): + } + }() + type outcome struct { + event *SessionEvent + err error + } + completion := make(chan outcome, 1) + var mu sync.Mutex + var messageID string + var admitted, started, completed bool + var pending []SessionEvent + var final *SessionEvent + finish := func(result outcome) { + if !completed { + completed = true + completion <- result + } + } + process := func(event SessionEvent) { + if completed || (event.AgentID != nil && *event.AgentID != "") { + return + } + switch data := event.Data.(type) { + case *UserMessageData: + if data.MessageID != nil && *data.MessageID == messageID { + started = true + } + case *AssistantMessageData: + if data.OriginatingMessageID != nil && *data.OriginatingMessageID == messageID { + started = true + if len(data.ToolRequests) > 0 { + final = nil + } else { + copy := event + final = © + } + } + case *SessionIdleData: + if !started || (data.Mode != nil && *data.Mode == SessionModeAutopilot) { + return + } + if data.Aborted != nil && *data.Aborted { + finish(outcome{err: fmt.Errorf("session aborted before structured output completed")}) + } else if final == nil || strings.TrimSpace(final.Data.(*AssistantMessageData).Content) == "" { + finish(outcome{err: fmt.Errorf("run completed without a structured assistant response")}) + } else { + finish(outcome{event: final}) + } + case *SessionErrorData: + if started { + finish(outcome{err: fmt.Errorf("session error: %s", data.Message)}) + } + } + } + unsubscribe := s.On(func(event SessionEvent) { + switch event.Data.(type) { + case *UserMessageData, *AssistantMessageData, *SessionIdleData, *SessionErrorData: + default: + return + } + mu.Lock() + defer mu.Unlock() + if !admitted { + pending = append(pending, event) + } else { + process(event) + } + }) + defer unsubscribe() + id, err := s.Send(ctx, options) + if err != nil { + if cause := context.Cause(ctx); cause != nil { + return nil, fmt.Errorf("admitting structured output: %w", cause) + } + return nil, err + } + mu.Lock() + messageID, admitted = id, true + for _, event := range pending { + process(event) + } + pending = nil + mu.Unlock() + select { + case result := <-completion: + return result.event, result.err + case <-s.eventDone: + return nil, fmt.Errorf("session closed before structured output completed") + case <-ctx.Done(): + return nil, fmt.Errorf("waiting for structured output: %w", context.Cause(ctx)) + } +} diff --git a/go/structured_output_test.go b/go/structured_output_test.go new file mode 100644 index 0000000000..c5196868eb --- /dev/null +++ b/go/structured_output_test.go @@ -0,0 +1,154 @@ +package copilot + +import ( + "context" + "io" + "reflect" + "strings" + "testing" + "time" + + "github.com/github/copilot-sdk/go/internal/jsonrpc2" +) + +type structuredInventory struct { + Count int `json:"count"` + Color string `json:"color"` +} + +func structuredMessage(content, origin string) SessionEvent { + return SessionEvent{Data: &AssistantMessageData{ + MessageID: "assistant", OriginatingMessageID: ptr(origin), Content: content, + }} +} + +func TestStructuredOutputTypedSchemaAndCorrelation(t *testing.T) { + subagent := structuredMessage(`{"count":999,"color":"wrong"}`, "message-1") + subagent.AgentID = ptr("child") + events := []SessionEvent{ + {Data: &SessionIdleData{}}, + {Data: &SessionErrorData{Message: "before this run"}}, + {Data: &UserMessageData{MessageID: ptr("message-1")}}, + structuredMessage(`{"count":42,"color":"red"}`, "message-1"), + {Data: &SessionIdleData{Mode: ptr(SessionModeAutopilot)}}, + structuredMessage(`{"count":99,"color":"blue"}`, "message-1"), + subagent, + structuredMessage(`{"count":123,"color":"wrong"}`, "other"), + {Data: &SessionIdleData{}}, + } + params := captureSessionSendRequest(t, nil, events, true, func(session *Session) { + options := MessageOptions{Prompt: "inventory"} + result, err := SendAndWait[structuredInventory](t.Context(), session, options) + if err != nil || result != (structuredInventory{Count: 99, Color: "blue"}) { + t.Fatalf("unexpected typed result: %+v, %v", result, err) + } + if options.ResponseSchema != nil { + t.Fatal("mutated caller options") + } + if len(session.handlers) != 0 { + t.Fatal("structured wait leaked subscription") + } + }) + format := params["responseFormat"].(map[string]any) + contract := format["jsonSchema"].(map[string]any) + schema := contract["schema"].(map[string]any) + if format["type"] != "json_schema" || contract["strict"] != true || schema["type"] != "object" { + t.Fatalf("unexpected response format: %#v", format) + } + if schema["additionalProperties"] != false { + t.Fatalf("typed struct schema must be closed for strict output: %#v", schema) + } + properties := schema["properties"].(map[string]any) + if properties["count"].(map[string]any)["type"] != "integer" || properties["color"].(map[string]any)["type"] != "string" { + t.Fatalf("incorrect inferred schema: %#v", schema) + } +} + +func TestStructuredOutputRawSchemaUnchanged(t *testing.T) { + schema := map[string]any{"type": "object", "description": "unmodified"} + events := []SessionEvent{structuredMessage(`{"count":42}`, "message-1"), {Data: &SessionIdleData{}}} + params := captureMessageSourceRequest(t, nil, events, func(session *Session) { + event, err := session.SendAndWait(t.Context(), MessageOptions{Prompt: "inventory", ResponseSchema: schema}) + if err != nil || event.Data.(*AssistantMessageData).Content != `{"count":42}` { + t.Fatalf("unexpected raw result: %v, %v", event, err) + } + }) + got := params["responseFormat"].(map[string]any)["jsonSchema"].(map[string]any)["schema"] + if !reflect.DeepEqual(got, schema) { + t.Fatalf("schema changed: %#v", got) + } +} + +func TestStructuredOutputFailures(t *testing.T) { + for _, tc := range []struct { + name string + events []SessionEvent + rpcError *jsonrpc2.Error + want string + }{ + {"missing", []SessionEvent{{Data: &UserMessageData{MessageID: ptr("message-1")}}, {Data: &SessionIdleData{}}}, nil, "without a structured"}, + {"blank", []SessionEvent{structuredMessage(" ", "message-1"), {Data: &SessionIdleData{}}}, nil, "without a structured"}, + {"aborted", []SessionEvent{structuredMessage(`{"count":1}`, "message-1"), {Data: &SessionIdleData{Aborted: ptr(true)}}}, nil, "aborted"}, + {"session error", []SessionEvent{{Data: &UserMessageData{MessageID: ptr("message-1")}}, {Data: &SessionErrorData{Message: "provider failed"}}}, nil, "provider failed"}, + {"admission", nil, &jsonrpc2.Error{Code: -32602, Message: "invalid schema"}, "invalid schema"}, + {"null", []SessionEvent{structuredMessage("null", "message-1"), {Data: &SessionIdleData{}}}, nil, "JSON null"}, + {"invalid json", []SessionEvent{structuredMessage("not JSON", "message-1"), {Data: &SessionIdleData{}}}, nil, "decode structured"}, + {"wrong type", []SessionEvent{structuredMessage(`{"count":"bad"}`, "message-1"), {Data: &SessionIdleData{}}}, nil, "decode structured"}, + } { + t.Run(tc.name, func(t *testing.T) { + captureMessageSourceRequest(t, tc.rpcError, tc.events, func(session *Session) { + ctx, cancel := context.WithTimeout(t.Context(), time.Second) + defer cancel() + _, err := SendAndWait[structuredInventory](ctx, session, MessageOptions{Prompt: "inventory"}) + if err == nil || !strings.Contains(err.Error(), tc.want) { + t.Fatalf("expected %q, got %v", tc.want, err) + } + if len(session.handlers) != 0 { + t.Fatal("leaked subscription") + } + }) + }) + } +} + +func TestStructuredOutputTypedRejectsConflictingOptions(t *testing.T) { + for _, options := range []MessageOptions{{Mode: "immediate"}, {ResponseSchema: map[string]any{}}} { + _, err := SendAndWait[structuredInventory](t.Context(), &Session{}, options) + if err == nil || !strings.Contains(err.Error(), "cannot specify") { + t.Fatalf("expected pre-admission rejection, got %v", err) + } + } +} + +func TestStructuredOutputDisconnectDuringAdmission(t *testing.T) { + stdinR, stdinW := io.Pipe() + stdoutR, stdoutW := io.Pipe() + defer stdinR.Close() + defer stdinW.Close() + defer stdoutR.Close() + defer stdoutW.Close() + client := jsonrpc2.NewClient(stdinW, stdoutR) + client.Start() + defer client.Stop() + session := newSession("session-1", client, "", false) + defer session.stopEventProcessing() + ctx, cancel := context.WithTimeout(t.Context(), time.Second) + defer cancel() + waiting := make(chan error, 1) + go func() { + _, err := SendAndWait[structuredInventory](ctx, session, MessageOptions{Prompt: "inventory"}) + waiting <- err + }() + if _, err := readTestJSONRPCFrame(stdinR); err != nil { + t.Fatal(err) + } + session.stopEventProcessing() + select { + case err := <-waiting: + if err == nil || !strings.Contains(err.Error(), "session closed") { + t.Fatalf("expected session closure while admission pending, got %v", err) + } + case <-ctx.Done(): + t.Fatal("session closure did not interrupt admission") + } +} diff --git a/go/types.go b/go/types.go index 7bc5bfb9af..cfbf5b5e8f 100644 --- a/go/types.go +++ b/go/types.go @@ -2455,6 +2455,10 @@ func MessageSourceAgent(id string) MessageSource { // MessageOptions configures a message to send type MessageOptions struct { + // ResponseSchema is a per-run JSON Schema. Independent sends do not inherit it. + // Immediate steering inherits the active format and must not specify its own. + // Use RPC.Send for schema name, description and strictness options. + ResponseSchema map[string]any // Prompt is the message to send Prompt string // Source identifies the message origin independently of Mode and AgentMode. @@ -2938,16 +2942,17 @@ type sessionAbortRequest struct { } type sessionSendRequest struct { - SessionID string `json:"sessionId"` - Prompt string `json:"prompt"` - Source MessageSource `json:"source,omitempty"` - DisplayPrompt string `json:"displayPrompt,omitempty"` - Attachments []Attachment `json:"attachments,omitempty"` - Mode string `json:"mode,omitempty"` - AgentMode AgentMode `json:"agentMode,omitempty"` - Traceparent string `json:"traceparent,omitempty"` - Tracestate string `json:"tracestate,omitempty"` - RequestHeaders map[string]string `json:"requestHeaders,omitempty"` + ResponseFormat *rpc.ResponseFormat `json:"responseFormat,omitempty"` + SessionID string `json:"sessionId"` + Prompt string `json:"prompt"` + Source MessageSource `json:"source,omitempty"` + DisplayPrompt string `json:"displayPrompt,omitempty"` + Attachments []Attachment `json:"attachments,omitempty"` + Mode string `json:"mode,omitempty"` + AgentMode AgentMode `json:"agentMode,omitempty"` + Traceparent string `json:"traceparent,omitempty"` + Tracestate string `json:"tracestate,omitempty"` + RequestHeaders map[string]string `json:"requestHeaders,omitempty"` } // sessionSendResponse is the response from session.send diff --git a/java/README.md b/java/README.md index af4cf4973f..153f838346 100644 --- a/java/README.md +++ b/java/README.md @@ -231,6 +231,41 @@ Agent sources serialize as `agent-`. Pass the agent ID without adding a prefix. The SDK preserves its case and whitespace and rejects null IDs. `sendAndWait` accepts the same source values as `send`. +## Structured output (experimental) + +Annotate a result record or class using the same compile-time schema-generation +approach as `@CopilotTool`. No additional schema dependency is needed: + +```java +@CopilotResponse +public record Inventory(int count, String color) {} + +Inventory inventory = session.sendAndWait( + "Call get_inventory, then report the widget count and color.", + Inventory.class +).get(); +``` + +Enable annotation processing with `CopilotResponseProcessor` (automatically +discoverable alongside the SDK's existing processors), and opt in to experimental +APIs as described below. The processor reuses the custom-tool `SchemaGenerator`, +closing record/class objects for strict output. Its existing type-mapping +limitations apply: custom Jackson naming/converters and recursive types need an explicit schema. +Provider restrictions, including optional-field and dictionary restrictions, +still apply. Jackson deserialization is not full JSON Schema validation. + +For an explicit schema, use `new MessageOptions().setPrompt(...).setResponseSchema(schema)` +with `send` or `sendAndWait`; schema-bearing `sendAndWait` returns the ordinary +message event. Typed overloads accept message options and a timeout, clone the +options, and reject explicit schemas and immediate delivery. + +Schemas apply to one run, including tools, steering, and stop-hook corrections; +independent sends and subagents do not inherit them. Streaming stays text. +Structured waits return the last correlated root message without tool requests at +non-autopilot idle. Concurrent waits keep their own results; queued work can delay +idle. Aborts, session errors after the run starts, and missing final output fail. +Cancellation and timeout stop waiting without aborting the agent. + ## Permission Handling `PermissionHandler.APPROVE_ALL` approves requests when managed settings are disabled. When `enableManagedSettings` is true, it completes exceptionally. Custom handlers can inspect `request.getManagedApprovalRequired()` for human-facing confirmation logic. diff --git a/java/sdk/pom.xml b/java/sdk/pom.xml index 8006873d93..96eef99522 100644 --- a/java/sdk/pom.xml +++ b/java/sdk/pom.xml @@ -153,6 +153,21 @@ none + + + default-testCompile + + full + + com.github.copilot.tool.CopilotResponseProcessor + + + -processorpath + ${project.build.outputDirectory} + + + + org.apache.maven.plugins diff --git a/java/sdk/src/main/java/com/github/copilot/CopilotResponse.java b/java/sdk/src/main/java/com/github/copilot/CopilotResponse.java new file mode 100644 index 0000000000..fd7b945cb4 --- /dev/null +++ b/java/sdk/src/main/java/com/github/copilot/CopilotResponse.java @@ -0,0 +1,23 @@ +/*--------------------------------------------------------------------------------------------- + * Copyright (c) Microsoft Corporation. All rights reserved. + *--------------------------------------------------------------------------------------------*/ + +package com.github.copilot; + +import java.lang.annotation.ElementType; +import java.lang.annotation.Retention; +import java.lang.annotation.RetentionPolicy; +import java.lang.annotation.Target; + +/** + * Generates a response schema using the same annotation processor and type + * mappings as {@code @CopilotTool}. Annotate a concrete record or class, then + * pass its class to {@link CopilotSession#sendAndWait(String, Class)}. Custom + * Jackson naming/converter schemas should instead be supplied explicitly via + * {@link com.github.copilot.rpc.MessageOptions#setResponseSchema(java.util.Map)}. + */ +@Target(ElementType.TYPE) +@Retention(RetentionPolicy.SOURCE) +@CopilotExperimental +public @interface CopilotResponse { +} diff --git a/java/sdk/src/main/java/com/github/copilot/CopilotSession.java b/java/sdk/src/main/java/com/github/copilot/CopilotSession.java index e0aab22e1f..38aafe801b 100644 --- a/java/sdk/src/main/java/com/github/copilot/CopilotSession.java +++ b/java/sdk/src/main/java/com/github/copilot/CopilotSession.java @@ -165,6 +165,7 @@ public final class CopilotSession implements AutoCloseable { private static final Logger LOG = Logger.getLogger(CopilotSession.class.getName()); private static final ObjectMapper MAPPER = JsonRpcClient.getObjectMapper(); + private final java.util.Set> structuredWaits = ConcurrentHashMap.newKeySet(); /** * Fixed name of the runtime's built-in tool-search tool. A client can replace @@ -565,6 +566,10 @@ public CompletableFuture send(MessageOptions options) { request.setAgentMode(options.getAgentMode()); request.setRequestHeaders(options.getRequestHeaders()); request.setDisplayPrompt(options.getDisplayPrompt()); + if (options.getResponseSchema() != null) { + request.setResponseFormat(Map.of("type", "json_schema", "jsonSchema", + Map.of("name", "response", "strict", true, "schema", options.getResponseSchema()))); + } return rpc.invoke("session.send", request, SendMessageResponse.class).thenApply(SendMessageResponse::messageId); } @@ -597,6 +602,9 @@ public CompletableFuture send(MessageOptions options) { */ public CompletableFuture sendAndWait(MessageOptions options, long timeoutMs) { ensureNotTerminated(); + if (options.getResponseSchema() != null) { + return sendAndWaitStructured(options, timeoutMs); + } long totalNanos = System.nanoTime(); var future = new CompletableFuture(); var lastAssistantMessage = new AtomicReference(); @@ -729,6 +737,176 @@ public CompletableFuture sendAndWait(MessageOptions optio return sendAndWait(options, 60000); } + /** + * Sends a prompt using the schema generated for a {@link CopilotResponse} type. + * + * @param + * the result type + * @param prompt + * the prompt + * @param responseType + * the annotated result class + * @return the deserialized, non-null final response + */ + @CopilotExperimental + public CompletableFuture sendAndWait(String prompt, Class responseType) { + return sendAndWait(new MessageOptions().setPrompt(prompt), responseType, 60000); + } + + /** + * Sends a message using its result type's generated schema. + * + * @param + * the result type + * @param options + * the message options, without an explicit schema or immediate + * delivery + * @param responseType + * the annotated result class + * @return the deserialized final response + */ + @CopilotExperimental + public CompletableFuture sendAndWait(MessageOptions options, Class responseType) { + return sendAndWait(options, responseType, 60000); + } + + /** + * Infers a schema using the custom-tool annotation processor, then deserializes + * the last correlated root message without tool requests at non-autopilot idle. + * Provider schema restrictions apply. Jackson deserialization is not full JSON + * Schema validation. Cancellation stops waiting without aborting agent work. + * + * @param + * the result type + * @param options + * the message options; not mutated + * @param responseType + * the annotated result class + * @param timeoutMs + * the wait timeout, or nonpositive for no timeout + * @return the non-null deserialized response + */ + @CopilotExperimental + public CompletableFuture sendAndWait(MessageOptions options, Class responseType, long timeoutMs) { + ensureNotTerminated(); + if (options.getResponseSchema() != null || "immediate".equals(options.getMode())) { + throw new IllegalArgumentException("Typed output cannot specify a response schema or immediate delivery"); + } + var message = options.clone().setResponseSchema(ResponseSchemas.forType(responseType)); + var waiting = sendAndWaitStructured(message, timeoutMs); + CompletableFuture result = waiting.thenApply(event -> { + try { + T value = MAPPER.readerFor(responseType) + .with(com.fasterxml.jackson.databind.DeserializationFeature.FAIL_ON_TRAILING_TOKENS) + .readValue(event.getData().content()); + if (value == null) + throw new IllegalStateException("Structured response was JSON null, not a result"); + return value; + } catch (com.fasterxml.jackson.core.JsonProcessingException e) { + throw new java.util.concurrent.CompletionException(e); + } + }); + result.whenComplete((value, error) -> { + if (result.isCancelled()) + waiting.cancel(true); + }); + return result; + } + + private CompletableFuture sendAndWaitStructured(MessageOptions options, long timeoutMs) { + var result = new CompletableFuture(); + class State { + String messageId; + boolean started; + AssistantMessageEvent last; + final List pending = new ArrayList<>(); + + synchronized void event(SessionEvent event) { + if (result.isDone() || (event.getAgentId() != null && !event.getAgentId().isEmpty())) + return; + if (messageId == null) { + pending.add(event); + return; + } + if (event instanceof com.github.copilot.generated.UserMessageEvent user + && messageId.equals(user.getData().messageId())) { + started = true; + } else if (event instanceof AssistantMessageEvent assistant + && messageId.equals(assistant.getData().originatingMessageId())) { + started = true; + last = assistant.getData().toolRequests() != null && !assistant.getData().toolRequests().isEmpty() + ? null + : assistant; + } else if (event instanceof SessionIdleEvent idle && started + && idle.getData().mode() != SessionMode.AUTOPILOT) { + if (Boolean.TRUE.equals(idle.getData().aborted())) { + result.completeExceptionally( + new IllegalStateException("Session aborted before structured output completed")); + } else if (last == null || last.getData().content().isBlank()) { + result.completeExceptionally( + new IllegalStateException("Run completed without a structured assistant response")); + } else { + result.complete(last); + } + } else if (event instanceof SessionErrorEvent error && started) { + result.completeExceptionally( + new IllegalStateException("Session error: " + error.getData().message())); + } + } + } + var state = new State(); + Closeable subscription = on(event -> { + if (event instanceof AssistantMessageEvent || event instanceof SessionIdleEvent + || event instanceof SessionErrorEvent + || event instanceof com.github.copilot.generated.UserMessageEvent) { + state.event(event); + } + }); + structuredWaits.add(result); + ScheduledFuture timer; + try { + timer = timeoutMs > 0 + ? timeoutScheduler.schedule( + () -> result.completeExceptionally( + new TimeoutException("Structured output timed out after " + timeoutMs + "ms")), + timeoutMs, TimeUnit.MILLISECONDS) + : null; + } catch (RejectedExecutionException e) { + result.completeExceptionally(e); + timer = null; + } + final ScheduledFuture timeout = timer; + result.whenComplete((value, error) -> { + structuredWaits.remove(result); + if (timeout != null) + timeout.cancel(false); + try { + subscription.close(); + } catch (IOException e) { + LOG.log(Level.SEVERE, "Error closing structured output subscription", e); + } + }); + if (!result.isDone()) { + try { + send(options).whenComplete((id, error) -> { + if (error != null) { + result.completeExceptionally(error); + } else { + synchronized (state) { + state.messageId = id; + for (var event : state.pending) + state.event(event); + state.pending.clear(); + } + } + }); + } catch (IllegalStateException | IllegalArgumentException e) { + result.completeExceptionally(e); + } + } + return result; + } + /** * Registers a callback for all session events. *

@@ -2527,6 +2705,8 @@ public void close() { isTerminated = true; } + structuredWaits.forEach(wait -> wait + .completeExceptionally(new IllegalStateException("Session closed before structured output completed"))); cancelPendingExternalTools(); timeoutScheduler.shutdownNow(); releaseGitHubTokenProviderRegistration(); diff --git a/java/sdk/src/main/java/com/github/copilot/ResponseSchemas.java b/java/sdk/src/main/java/com/github/copilot/ResponseSchemas.java new file mode 100644 index 0000000000..47423b19b2 --- /dev/null +++ b/java/sdk/src/main/java/com/github/copilot/ResponseSchemas.java @@ -0,0 +1,26 @@ +/*--------------------------------------------------------------------------------------------- + * Copyright (c) Microsoft Corporation. All rights reserved. + *--------------------------------------------------------------------------------------------*/ + +package com.github.copilot; + +import java.lang.reflect.InvocationTargetException; +import java.util.Map; + +final class ResponseSchemas { + private ResponseSchemas() { + } + + @SuppressWarnings("unchecked") + static Map forType(Class type) { + try { + Class metadata = Class.forName(type.getName() + "$$CopilotResponseMeta", true, type.getClassLoader()); + return (Map) metadata.getMethod("schema").invoke(null); + } catch (ClassNotFoundException e) { + throw new IllegalArgumentException("No response schema for " + type.getName() + + ". Annotate the type with @CopilotResponse and enable CopilotResponseProcessor.", e); + } catch (NoSuchMethodException | IllegalAccessException | InvocationTargetException e) { + throw new IllegalStateException("Cannot load response schema for " + type.getName(), e); + } + } +} diff --git a/java/sdk/src/main/java/com/github/copilot/rpc/MessageOptions.java b/java/sdk/src/main/java/com/github/copilot/rpc/MessageOptions.java index 10458736e1..b06264a0c4 100644 --- a/java/sdk/src/main/java/com/github/copilot/rpc/MessageOptions.java +++ b/java/sdk/src/main/java/com/github/copilot/rpc/MessageOptions.java @@ -49,6 +49,31 @@ public class MessageOptions { private AgentMode agentMode; private Map requestHeaders; private String displayPrompt; + private Map responseSchema; + + /** + * Gets this run's output schema. + * + * @return the schema, or null when unformatted + */ + @com.github.copilot.CopilotExperimental + public Map getResponseSchema() { + return responseSchema == null ? null : Collections.unmodifiableMap(responseSchema); + } + + /** + * Requests provider-native JSON Schema output for this run. Independent sends + * do not inherit it. Immediate steering cannot specify a schema. + * + * @param responseSchema + * the schema, passed unchanged to the provider + * @return this options instance + */ + @com.github.copilot.CopilotExperimental + public MessageOptions setResponseSchema(Map responseSchema) { + this.responseSchema = responseSchema == null ? null : new HashMap<>(responseSchema); + return this; + } /** * Gets the message prompt. @@ -246,6 +271,7 @@ public MessageOptions clone() { copy.agentMode = this.agentMode; copy.requestHeaders = this.requestHeaders != null ? new HashMap<>(this.requestHeaders) : null; copy.displayPrompt = this.displayPrompt; + copy.responseSchema = this.responseSchema == null ? null : new HashMap<>(this.responseSchema); return copy; } diff --git a/java/sdk/src/main/java/com/github/copilot/rpc/SendMessageRequest.java b/java/sdk/src/main/java/com/github/copilot/rpc/SendMessageRequest.java index 97cfc76f9d..fb27324831 100644 --- a/java/sdk/src/main/java/com/github/copilot/rpc/SendMessageRequest.java +++ b/java/sdk/src/main/java/com/github/copilot/rpc/SendMessageRequest.java @@ -49,6 +49,19 @@ public final class SendMessageRequest { @JsonProperty("displayPrompt") private String displayPrompt; + @JsonProperty("responseFormat") + private Map responseFormat; + + /** Gets the output format. @return the provider-native format */ + public Map getResponseFormat() { + return responseFormat; + } + + /** Sets the output format. @param responseFormat the provider-native format */ + public void setResponseFormat(Map responseFormat) { + this.responseFormat = responseFormat; + } + /** Gets the session ID. @return the session ID */ public String getSessionId() { return sessionId; diff --git a/java/sdk/src/main/java/com/github/copilot/tool/CopilotResponseProcessor.java b/java/sdk/src/main/java/com/github/copilot/tool/CopilotResponseProcessor.java new file mode 100644 index 0000000000..fa98980165 --- /dev/null +++ b/java/sdk/src/main/java/com/github/copilot/tool/CopilotResponseProcessor.java @@ -0,0 +1,71 @@ +/*--------------------------------------------------------------------------------------------- + * Copyright (c) Microsoft Corporation. All rights reserved. + *--------------------------------------------------------------------------------------------*/ + +package com.github.copilot.tool; + +import java.io.IOException; +import java.io.PrintWriter; +import java.util.Set; + +import javax.annotation.processing.AbstractProcessor; +import javax.annotation.processing.RoundEnvironment; +import javax.annotation.processing.SupportedAnnotationTypes; +import javax.annotation.processing.SupportedSourceVersion; +import javax.lang.model.SourceVersion; +import javax.lang.model.element.Element; +import javax.lang.model.element.ElementKind; +import javax.lang.model.element.Modifier; +import javax.lang.model.element.TypeElement; +import javax.tools.Diagnostic; + +import com.github.copilot.CopilotExperimental; +import com.github.copilot.CopilotResponse; + +/** + * Generates response metadata using the existing custom-tool schema generator. + */ +@SupportedAnnotationTypes("com.github.copilot.CopilotResponse") +@SupportedSourceVersion(SourceVersion.RELEASE_17) +@CopilotExperimental +public class CopilotResponseProcessor extends AbstractProcessor { + @Override + public boolean process(Set annotations, RoundEnvironment roundEnv) { + for (Element element : roundEnv.getElementsAnnotatedWith(CopilotResponse.class)) { + if (!(element instanceof TypeElement type) + || (type.getKind() != ElementKind.CLASS && type.getKind() != ElementKind.RECORD) + || type.getModifiers().contains(Modifier.PRIVATE) || !type.getTypeParameters().isEmpty()) { + processingEnv.getMessager().printMessage(Diagnostic.Kind.ERROR, + "@CopilotResponse requires an accessible, non-generic class or record", element); + continue; + } + String binaryName = processingEnv.getElementUtils().getBinaryName(type).toString(); + String packageName = processingEnv.getElementUtils().getPackageOf(type).getQualifiedName().toString(); + String metadataName = binaryName.substring(packageName.isEmpty() ? 0 : packageName.length() + 1) + + "$$CopilotResponseMeta"; + String qualifiedName = packageName.isEmpty() ? metadataName : packageName + "." + metadataName; + String schema; + try { + schema = new SchemaGenerator(true).generateSchemaSource(type.asType(), processingEnv.getTypeUtils(), + processingEnv.getElementUtils()); + } catch (IllegalArgumentException e) { + processingEnv.getMessager().printMessage(Diagnostic.Kind.ERROR, e.getMessage(), type); + continue; + } + try (var writer = new PrintWriter( + processingEnv.getFiler().createSourceFile(qualifiedName, type).openWriter())) { + if (!packageName.isEmpty()) + writer.println("package " + packageName + ";"); + writer.println("import java.util.Map;"); + writer.println("import java.util.List;"); + writer.println("public final class " + metadataName + " {"); + writer.println("public static Map schema() { return " + schema + "; }"); + writer.println("}"); + } catch (IOException e) { + processingEnv.getMessager().printMessage(Diagnostic.Kind.ERROR, + "Cannot generate response schema: " + e.getMessage(), type); + } + } + return true; + } +} diff --git a/java/sdk/src/main/java/com/github/copilot/tool/SchemaGenerator.java b/java/sdk/src/main/java/com/github/copilot/tool/SchemaGenerator.java index 59336a1e02..0cfa2e512a 100644 --- a/java/sdk/src/main/java/com/github/copilot/tool/SchemaGenerator.java +++ b/java/sdk/src/main/java/com/github/copilot/tool/SchemaGenerator.java @@ -5,7 +5,9 @@ package com.github.copilot.tool; import java.util.ArrayList; +import java.util.HashSet; import java.util.List; +import java.util.Set; import java.util.stream.Collectors; import javax.lang.model.element.Element; @@ -35,6 +37,17 @@ */ @CopilotExperimental public class SchemaGenerator { + private final boolean closedObjects; + private final Set expandingResponseTypes = new HashSet<>(); + + /** Creates a schema generator for custom tools. */ + public SchemaGenerator() { + this(false); + } + + SchemaGenerator(boolean closedObjects) { + this.closedObjects = closedObjects; + } /** * Given a {@link TypeMirror} from the annotation processing environment, @@ -100,7 +113,7 @@ public String generateParametersSchemaSource(List par String properties = "Map.ofEntries(" + String.join(", ", propertyEntries) + ")"; String required = "List.of(" + String.join(", ", requiredNames) + ")"; - return "Map.of(\"type\", \"object\", \"properties\", " + properties + ", \"required\", " + required + ")"; + return objectSchema(properties, required); } private String generateSchema(TypeMirror type, Types typeUtils, Elements elementUtils) { @@ -119,6 +132,18 @@ private String generateSchema(TypeMirror type, Types typeUtils, Elements element // Handle declared types (classes, interfaces, enums, records) if (type.getKind() == TypeKind.DECLARED) { + if (closedObjects) { + String name = type.toString(); + if (!expandingResponseTypes.add(name)) { + throw new IllegalArgumentException( + "Recursive response types require an explicit JSON Schema: " + name); + } + try { + return generateDeclaredTypeSchema((DeclaredType) type, typeUtils, elementUtils); + } finally { + expandingResponseTypes.remove(name); + } + } return generateDeclaredTypeSchema((DeclaredType) type, typeUtils, elementUtils); } @@ -294,7 +319,7 @@ private String generateRecordSchema(TypeElement typeElement, Types typeUtils, El String properties = "Map.ofEntries(" + String.join(", ", propertyEntries) + ")"; String required = "List.of(" + String.join(", ", requiredNames) + ")"; - return "Map.of(\"type\", \"object\", \"properties\", " + properties + ", \"required\", " + required + ")"; + return objectSchema(properties, required); } private String generateClassSchema(TypeElement typeElement, Types typeUtils, Elements elementUtils) { @@ -325,13 +350,18 @@ private String generateClassSchema(TypeElement typeElement, Types typeUtils, Ele } if (propertyEntries.isEmpty()) { - return "Map.of(\"type\", \"object\")"; + return closedObjects ? objectSchema("Map.of()", "List.of()") : "Map.of(\"type\", \"object\")"; } String properties = "Map.ofEntries(" + String.join(", ", propertyEntries) + ")"; String required = "List.of(" + String.join(", ", requiredNames) + ")"; - return "Map.of(\"type\", \"object\", \"properties\", " + properties + ", \"required\", " + required + ")"; + return objectSchema(properties, required); + } + + private String objectSchema(String properties, String required) { + return "Map.of(\"type\", \"object\", \"properties\", " + properties + ", \"required\", " + required + + (closedObjects ? ", \"additionalProperties\", false" : "") + ")"; } private String generateSealedSchema(TypeElement typeElement, Types typeUtils, Elements elementUtils) { diff --git a/java/sdk/src/main/java/module-info.java b/java/sdk/src/main/java/module-info.java index 8bc2dbd55c..1579758a76 100644 --- a/java/sdk/src/main/java/module-info.java +++ b/java/sdk/src/main/java/module-info.java @@ -28,6 +28,6 @@ opens com.github.copilot.rpc to com.fasterxml.jackson.databind; opens com.github.copilot.ffi to com.sun.jna; - provides javax.annotation.processing.Processor - with com.github.copilot.CopilotExperimentalProcessor, com.github.copilot.tool.CopilotToolProcessor; + provides javax.annotation.processing.Processor with com.github.copilot.CopilotExperimentalProcessor, + com.github.copilot.tool.CopilotToolProcessor, com.github.copilot.tool.CopilotResponseProcessor; } diff --git a/java/sdk/src/main/resources/META-INF/services/javax.annotation.processing.Processor b/java/sdk/src/main/resources/META-INF/services/javax.annotation.processing.Processor index 3b2e17d2f9..e5eba6f25b 100644 --- a/java/sdk/src/main/resources/META-INF/services/javax.annotation.processing.Processor +++ b/java/sdk/src/main/resources/META-INF/services/javax.annotation.processing.Processor @@ -1,2 +1,3 @@ com.github.copilot.CopilotExperimentalProcessor com.github.copilot.tool.CopilotToolProcessor +com.github.copilot.tool.CopilotResponseProcessor diff --git a/java/sdk/src/test/java/com/github/copilot/StructuredOutputE2ETest.java b/java/sdk/src/test/java/com/github/copilot/StructuredOutputE2ETest.java new file mode 100644 index 0000000000..c38f343fb8 --- /dev/null +++ b/java/sdk/src/test/java/com/github/copilot/StructuredOutputE2ETest.java @@ -0,0 +1,385 @@ +/*--------------------------------------------------------------------------------------------- + * Copyright (c) Microsoft Corporation. All rights reserved. + *--------------------------------------------------------------------------------------------*/ + +package com.github.copilot; + +import static org.junit.jupiter.api.Assertions.*; + +import java.util.List; +import java.util.Map; +import java.util.concurrent.CompletableFuture; +import java.util.concurrent.CopyOnWriteArrayList; +import java.util.concurrent.TimeUnit; +import java.util.concurrent.atomic.AtomicInteger; + +import org.junit.jupiter.api.AfterAll; +import org.junit.jupiter.api.BeforeAll; +import org.junit.jupiter.api.Test; + +import com.github.copilot.generated.AssistantMessageEvent; +import com.github.copilot.generated.AssistantMessageDeltaEvent; +import com.github.copilot.generated.SessionErrorEvent; +import com.github.copilot.generated.SessionIdleEvent; +import com.github.copilot.generated.UserMessageEvent; +import com.github.copilot.generated.rpc.SessionSendMessagesParams; +import com.github.copilot.rpc.AgentStopHookOutput; +import com.github.copilot.rpc.MessageOptions; +import com.github.copilot.rpc.PermissionHandler; +import com.github.copilot.rpc.ProviderConfig; +import com.github.copilot.rpc.SessionConfig; +import com.github.copilot.rpc.SessionHooks; +import com.github.copilot.rpc.ToolDefinition; + +class StructuredOutputE2ETest { + @CopilotResponse + public record Inventory(int count, String color) { + } + @CopilotResponse + public record Answer(int answer) { + } + @CopilotResponse + public record ToolAnswer(int answer, String contract) { + } + @CopilotResponse + public record First(int first) { + } + @CopilotResponse + public record Second(int second) { + } + + private static E2ETestContext ctx; + + @BeforeAll + static void setup() throws Exception { + ctx = E2ETestContext.create(); + } + + @AfterAll + static void teardown() throws Exception { + if (ctx != null) + ctx.close(); + } + + private SessionConfig config() { + return new SessionConfig().setModel("gpt-4.1").setAvailableTools(List.of()) + .setOnPermissionRequest(PermissionHandler.APPROVE_ALL).setProvider( + new ProviderConfig().setType("openai").setWireApi("completions").setBaseUrl(ctx.getProxyUrl()) + .setModelId("gpt-4.1").setWireModel("gpt-4.1").setApiKey("fake-token-for-e2e-tests") + .setHeaders(Map.of("Copilot-Integration-Id", "copilot-developer-cli", + "Copilot-Harness-Id", "copilot-sdk", "X-GitHub-Api-Version", "2026-08-01"))); + } + + @Test + void infersTypedResultAfterCustomTool() throws Exception { + ctx.configureForTest("structured_output", "infers_typed_result_after_custom_tool"); + var calls = new AtomicInteger(); + var tool = ToolDefinition.from("get_inventory", "Get the current widget inventory.", () -> { + calls.incrementAndGet(); + return "The inventory contains 42 red widgets."; + }); + try (var client = ctx.createClient(); + var session = client.createSession(config().setTools(List.of(tool)).setStreaming(true)).get()) { + var deltas = new AtomicInteger(); + try (var subscription = session.on(event -> { + if (event instanceof AssistantMessageDeltaEvent) + deltas.incrementAndGet(); + })) { + var result = session + .sendAndWait("Call get_inventory, then report the widget count and color.", Inventory.class) + .get(30, TimeUnit.SECONDS); + assertEquals(new Inventory(42, "red"), result); + assertTrue(calls.get() > 0); + assertTrue(deltas.get() > 0, "Typed wait must preserve streaming text updates"); + } + var ordinary = session + .sendAndWait( + new MessageOptions().setPrompt("Now reply with exactly the plain text HELLO, not JSON.")) + .get(30, TimeUnit.SECONDS); + assertEquals("HELLO", ordinary.getData().content().trim()); + var exchanges = ctx.getExchanges(); + assertTrue(exchanges.size() >= 3); + for (var exchange : exchanges.subList(0, exchanges.size() - 1)) { + var request = JsonRpcClient.getObjectMapper().valueToTree(exchange).get("request"); + var format = request.get("response_format"); + assertEquals("json_schema", format.get("type").asText()); + assertTrue(format.get("json_schema").get("strict").asBoolean()); + assertEquals(JsonRpcClient.getObjectMapper().valueToTree(ResponseSchemas.forType(Inventory.class)), + format.get("json_schema").get("schema")); + } + assertFalse(JsonRpcClient.getObjectMapper().valueToTree(exchanges.get(exchanges.size() - 1)).get("request") + .has("response_format")); + } + } + + @Test + void typedWaitReturnsStopHookCorrection() throws Exception { + ctx.configureForTest("structured_output", "typed_wait_returns_stop_hook_correction"); + var stops = new AtomicInteger(); + var hooks = new SessionHooks().setOnAgentStop((input, + invocation) -> CompletableFuture.completedFuture(stops.incrementAndGet() == 1 + ? new AgentStopHookOutput().setDecision("block") + .setReason("Correct the answer to 99, not 42. Do not use tools.") + : null)); + try (var client = ctx.createClient(); var session = client.createSession(config().setHooks(hooks)).get()) { + var replies = new CopyOnWriteArrayList(); + try (var subscription = session.on(event -> { + if (event instanceof AssistantMessageEvent message && event.getAgentId() == null) + replies.add(message); + })) { + var result = session.sendAndWait("What is 19 + 23? Do not use tools.", Answer.class).get(30, + TimeUnit.SECONDS); + assertEquals(99, result.answer()); + assertEquals(2, stops.get()); + assertEquals(2, replies.size()); + assertNotNull(replies.get(0).getData().originatingMessageId()); + assertFalse(replies.get(0).getData().originatingMessageId().isEmpty()); + assertEquals(replies.get(0).getData().originatingMessageId(), + replies.get(1).getData().originatingMessageId()); + assertEquals(42, JsonRpcClient.getObjectMapper() + .readValue(replies.get(0).getData().content(), Answer.class).answer()); + assertEquals(99, JsonRpcClient.getObjectMapper() + .readValue(replies.get(1).getData().content(), Answer.class).answer()); + } + } + } + + @Test + void sendSelectsCorrelatedResponseAfterIdle() throws Exception { + ctx.configureForTest("structured_output", "send_selects_correlated_response_after_idle"); + var entered = new CompletableFuture(); + var release = new CompletableFuture(); + var tool = ToolDefinition.from("read_inventory", "Read the current widget count and color.", + () -> "The inventory contains 42 red widgets."); + var hooks = new SessionHooks().setOnAgentStop((input, invocation) -> { + entered.complete(null); + return release.thenApply(ignored -> null); + }); + try (var client = ctx.createClient(); + var session = client.createSession(config().setTools(List.of(tool)).setHooks(hooks)).get()) { + var idle = new CompletableFuture(); + var replies = new CopyOnWriteArrayList(); + try (var subscription = session.on(event -> { + if (event.getAgentId() != null) + return; + if (event instanceof AssistantMessageEvent message) + replies.add(message); + else if (event instanceof SessionIdleEvent) + idle.complete(null); + else if (event instanceof SessionErrorEvent error) { + var failure = new IllegalStateException(error.getData().message()); + entered.completeExceptionally(failure); + idle.completeExceptionally(failure); + } + })) { + var origin = session.send(new MessageOptions() + .setPrompt("Call read_inventory once, then report the current widget count and color.") + .setResponseSchema(ResponseSchemas.forType(Inventory.class))).get(30, TimeUnit.SECONDS); + entered.get(30, TimeUnit.SECONDS); + assertFalse(idle.isDone(), "Idle must wait for the stop hook"); + release.complete(null); + idle.get(30, TimeUnit.SECONDS); + assertTrue(replies.size() >= 2); + var last = replies.get(replies.size() - 1); + assertEquals(origin, last.getData().originatingMessageId()); + assertTrue(last.getData().toolRequests() == null || last.getData().toolRequests().isEmpty()); + assertTrue(replies.stream().anyMatch( + reply -> reply.getData().toolRequests() != null && !reply.getData().toolRequests().isEmpty())); + assertEquals(new Inventory(42, "red"), + JsonRpcClient.getObjectMapper().readValue(last.getData().content(), Inventory.class)); + } finally { + release.complete(null); + } + } + } + + @Test + void rejectsInvalidFormatsBeforeAdmission() throws Exception { + ctx.initializeProxy(); + try (var client = ctx.createClient(); var session = client.createSession(config()).get()) { + var admitted = new AtomicInteger(); + try (var subscription = session.on(event -> { + if (event instanceof UserMessageEvent || event instanceof SessionErrorEvent) + admitted.incrementAndGet(); + })) { + var immediate = assertThrows(IllegalArgumentException.class, + () -> session.sendAndWait( + new MessageOptions().setPrompt("Must not be admitted").setMode("immediate"), + Answer.class).get(30, TimeUnit.SECONDS)); + assertTrue(immediate.getMessage().contains("immediate")); + var schema = Map.of("type", "object", "description", "x".repeat(32 * 1024 * 1024)); + var oversized = assertThrows(java.util.concurrent.ExecutionException.class, + () -> session.sendAndWait( + new MessageOptions().setPrompt("Must not be admitted").setResponseSchema(schema)) + .get(30, TimeUnit.SECONDS)); + assertTrue(oversized.getCause().getMessage().contains("32 MiB")); + var params = JsonRpcClient.getObjectMapper().convertValue( + Map.of("messages", List.of(), "responseFormat", + Map.of("type", "json_schema", "jsonSchema", + Map.of("name", "response", "schema", schema))), + SessionSendMessagesParams.class); + var batch = assertThrows(java.util.concurrent.ExecutionException.class, + () -> session.getRpc().sendMessages(params).get(30, TimeUnit.SECONDS)); + assertTrue(batch.getCause().getMessage().contains("32 MiB")); + assertTrue(session.getRpc().queue.pendingItems().get(30, TimeUnit.SECONDS).items().isEmpty()); + assertEquals(0, admitted.get()); + assertTrue(ctx.getExchanges().isEmpty()); + } + } + } + + @Test + void sendsExplicitSchemaForMessageAndBatch() throws Exception { + ctx.configureForTest("structured_output", "sends_explicit_schema_for_message_and_batch"); + try (var client = ctx.createClient(); var session = client.createSession(config()).get()) { + var idle = new CompletableFuture(); + var replies = new CopyOnWriteArrayList(); + try (var subscription = session.on(event -> { + if (event.getAgentId() != null) + return; + if (event instanceof AssistantMessageEvent message) + replies.add(message); + else if (event instanceof SessionIdleEvent) + idle.complete(null); + else if (event instanceof SessionErrorEvent error) + idle.completeExceptionally(new IllegalStateException(error.getData().message())); + })) { + var params = JsonRpcClient.getObjectMapper() + .convertValue( + Map.of("messages", + List.of(Map.of("prompt", "There are 42 red widgets in stock."), + Map.of("prompt", "Report the widget count and color.")), + "responseFormat", + Map.of("type", "json_schema", "jsonSchema", + Map.of("name", "inventory", "strict", true, "schema", + ResponseSchemas.forType(Inventory.class)))), + SessionSendMessagesParams.class); + var accepted = session.getRpc().sendMessages(params).get(30, TimeUnit.SECONDS); + idle.get(30, TimeUnit.SECONDS); + var finalMessage = replies.stream().filter( + message -> accepted.messageIds().get(1).equals(message.getData().originatingMessageId())) + .reduce((first, second) -> second).orElseThrow(); + assertEquals(new Inventory(42, "red"), + JsonRpcClient.getObjectMapper().readValue(finalMessage.getData().content(), Inventory.class)); + } + var raw = session.sendAndWait(new MessageOptions() + .setPrompt("The inventory now has 21 blue widgets. Report the new count and color.") + .setResponseSchema(ResponseSchemas.forType(Inventory.class))).get(30, TimeUnit.SECONDS); + assertEquals(new Inventory(21, "blue"), + JsonRpcClient.getObjectMapper().readValue(raw.getData().content(), Inventory.class)); + } + } + + @Test + void typedWaitReturnsStopHookCorrectionAfterTerminalTool() throws Exception { + ctx.configureForTest("structured_output", "typed_wait_returns_stop_hook_correction_after_terminal_tool"); + var calls = new AtomicInteger(); + var stops = new AtomicInteger(); + var tool = ToolDefinition.from("lookup_number", "Return the number needed for the calculation.", () -> { + calls.incrementAndGet(); + return 58; + }).isTerminal(true).skipPermission(true); + var hooks = new SessionHooks().setOnAgentStop((input, + invocation) -> CompletableFuture.completedFuture(stops.incrementAndGet() == 1 + ? new AgentStopHookOutput().setDecision("block") + .setReason("Correct the answer to 99, not 63. Do not use tools.") + : null)); + try (var client = ctx.createClient(); + var session = client.createSession(config().setTools(List.of(tool)).setHooks(hooks)).get()) { + var result = session.sendAndWait( + "Call lookup_number exactly once, then add 5 to the returned number. Do not guess its result.", + Answer.class).get(30, TimeUnit.SECONDS); + assertEquals(99, result.answer()); + assertEquals(1, calls.get()); + assertEquals(2, stops.get()); + } + } + + @Test + void typedWaitReturnsLateSteeringResponse() throws Exception { + ctx.configureForTest("structured_output", "typed_wait_returns_late_steering_response"); + var stops = new AtomicInteger(); + var current = new java.util.concurrent.atomic.AtomicReference(); + var hooks = new SessionHooks().setOnAgentStop((input, invocation) -> { + if (stops.incrementAndGet() == 1) { + return current.get().send(new MessageOptions().setPrompt("Change the answer to 99. Do not use tools.") + .setMode("immediate")).thenApply(id -> null); + } + return CompletableFuture.completedFuture(null); + }); + try (var client = ctx.createClient(); var session = client.createSession(config().setHooks(hooks)).get()) { + current.set(session); + assertEquals(99, session.sendAndWait("What is 19 + 23? Do not use tools.", Answer.class) + .get(30, TimeUnit.SECONDS).answer()); + assertEquals(2, stops.get()); + } + } + + @Test + void typedResultAfterTerminalToolAndSteering() throws Exception { + ctx.configureForTest("structured_output", "typed_result_after_terminal_tool_and_steering"); + var current = new java.util.concurrent.atomic.AtomicReference(); + var calls = new AtomicInteger(); + var tool = ToolDefinition.create("lookup_number", "Return the number needed for the calculation.", + Map.of("type", "object", "properties", Map.of(), "required", List.of()), invocation -> { + calls.incrementAndGet(); + return current.get() + .send(new MessageOptions() + .setPrompt("Continue with the original calculation. Do not call any more tools.") + .setMode("immediate")) + .thenApply(id -> 58); + }).isTerminal(true).skipPermission(true); + try (var client = ctx.createClient(); + var session = client.createSession(config().setTools(List.of(tool))).get()) { + current.set(session); + var result = session.sendAndWait( + "Call lookup_number exactly once, then add 5 to the returned number. Do not guess its result.", + ToolAnswer.class).get(30, TimeUnit.SECONDS); + assertEquals(new ToolAnswer(63, "typed_tool"), result); + assertEquals(1, calls.get()); + var exchanges = ctx.getExchanges(); + assertTrue(exchanges.size() >= 2); + for (var exchange : exchanges.subList(1, exchanges.size())) { + assertEquals("none", JsonRpcClient.getObjectMapper().valueToTree(exchange).get("request") + .get("tool_choice").asText()); + } + for (var exchange : exchanges) { + assertEquals(JsonRpcClient.getObjectMapper().valueToTree(ResponseSchemas.forType(ToolAnswer.class)), + JsonRpcClient.getObjectMapper().valueToTree(exchange).get("request").get("response_format") + .get("json_schema").get("schema")); + } + } + } + + @Test + void concurrentTypedSendsReturnTheirOwnResults() throws Exception { + ctx.configureForTest("structured_output", "concurrent_typed_sends_return_their_own_results"); + var entered = new CompletableFuture(); + var release = new CompletableFuture(); + var tool = ToolDefinition.create("first_number", "Get the number for the first question.", + Map.of("type", "object", "properties", Map.of(), "required", List.of()), invocation -> { + entered.complete(null); + return release; + }); + try (var client = ctx.createClient(); + var session = client.createSession(config().setTools(List.of(tool))).get()) { + var first = session.sendAndWait("Call first_number exactly once and report its returned number.", + First.class); + try { + CompletableFuture.anyOf(entered, first).get(30, TimeUnit.SECONDS); + assertTrue(entered.isDone(), "First run must enter its tool"); + var second = session.sendAndWait("What is 30 + 7? Do not use tools.", Second.class); + long deadline = System.nanoTime() + TimeUnit.SECONDS.toNanos(30); + while (session.getRpc().queue.pendingItems().get(30, TimeUnit.SECONDS).items().isEmpty()) { + assertTrue(System.nanoTime() < deadline, "Second run was not queued"); + Thread.sleep(10); + } + release.complete(42); + assertEquals(42, first.get(30, TimeUnit.SECONDS).first()); + assertEquals(37, second.get(30, TimeUnit.SECONDS).second()); + } finally { + release.complete(42); + } + } + } +} diff --git a/java/sdk/src/test/java/com/github/copilot/StructuredOutputTest.java b/java/sdk/src/test/java/com/github/copilot/StructuredOutputTest.java new file mode 100644 index 0000000000..b7aa6556bd --- /dev/null +++ b/java/sdk/src/test/java/com/github/copilot/StructuredOutputTest.java @@ -0,0 +1,186 @@ +/*--------------------------------------------------------------------------------------------- + * Copyright (c) Microsoft Corporation. All rights reserved. + *--------------------------------------------------------------------------------------------*/ + +package com.github.copilot; + +import static org.junit.jupiter.api.Assertions.*; +import static org.mockito.ArgumentMatchers.*; +import static org.mockito.Mockito.*; + +import java.util.List; +import java.util.Map; +import java.util.concurrent.CompletableFuture; +import java.util.concurrent.ExecutionException; +import java.util.concurrent.TimeUnit; + +import org.junit.jupiter.api.AfterEach; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.ValueSource; + +import com.github.copilot.generated.SessionEvent; +import com.github.copilot.rpc.MessageOptions; +import com.github.copilot.rpc.SendMessageRequest; +import com.github.copilot.rpc.SendMessageResponse; + +class StructuredOutputTest { + @CopilotResponse + public record Inventory(int count, String color) { + } + + @CopilotResponse + public record Nested(List items) { + } + + private JsonRpcClient rpc; + private CopilotSession session; + + @BeforeEach + void setup() { + rpc = mock(JsonRpcClient.class); + session = new CopilotSession("session-1", rpc, null); + when(rpc.invoke(eq("session.send"), any(), eq(SendMessageResponse.class))) + .thenReturn(CompletableFuture.completedFuture(new SendMessageResponse("user-1"))); + when(rpc.invoke(eq("session.detach"), any(), eq(CopilotSession.SessionDetachResponse.class))) + .thenReturn(CompletableFuture.completedFuture(new CopilotSession.SessionDetachResponse(true, null))); + } + + @AfterEach + void close() { + session.close(); + } + + private void emit(String type, Map data) throws Exception { + var event = JsonRpcClient.getObjectMapper().convertValue(Map.of("id", java.util.UUID.randomUUID().toString(), + "timestamp", "2026-01-01T00:00:00Z", "type", type, "data", data), SessionEvent.class); + var dispatch = CopilotSession.class.getDeclaredMethod("dispatchEvent", SessionEvent.class); + dispatch.setAccessible(true); + dispatch.invoke(session, event); + } + + private void assistant(String origin, String content) throws Exception { + emit("assistant.message", Map.of("messageId", "assistant", "originatingMessageId", origin, "content", content)); + } + + @Test + void infersClosedNestedSchemasUsingToolProcessor() { + var schema = JsonRpcClient.getObjectMapper().valueToTree(ResponseSchemas.forType(Nested.class)); + assertFalse(schema.path("additionalProperties").asBoolean(true)); + var inventory = schema.path("properties").path("items").path("items"); + assertEquals("integer", inventory.path("properties").path("count").path("type").asText()); + assertFalse(inventory.path("additionalProperties").asBoolean(true)); + assertEquals(2, inventory.path("required").size()); + } + + @Test + void buffersPreAckEventsAndDoesNotMutateOptions() throws Exception { + when(rpc.invoke(eq("session.send"), any(), eq(SendMessageResponse.class))).thenAnswer(invocation -> { + SendMessageRequest request = invocation.getArgument(1); + assertEquals(ResponseSchemas.forType(Inventory.class), + ((Map) request.getResponseFormat().get("jsonSchema")).get("schema")); + assistant("user-1", "{\"count\":42,\"color\":\"red\"}"); + emit("session.idle", Map.of()); + return CompletableFuture.completedFuture(new SendMessageResponse("user-1")); + }); + var options = new MessageOptions().setPrompt("inventory"); + assertEquals(new Inventory(42, "red"), session.sendAndWait(options, Inventory.class).get(1, TimeUnit.SECONDS)); + assertNull(options.getResponseSchema()); + } + + @Test + void ignoresUnrelatedMessagesAndWaitsForCorrectedFinalIdle() throws Exception { + var result = session.sendAndWait("inventory", Inventory.class); + emit("session.idle", Map.of()); + assistant("user-1", "{\"count\":42,\"color\":\"red\"}"); + emit("session.idle", Map.of("mode", "autopilot")); + assertFalse(result.isDone()); + assistant("user-1", "{\"count\":99,\"color\":\"blue\"}"); + assistant("other", "{\"count\":123,\"color\":\"wrong\"}"); + emit("session.idle", Map.of()); + assertEquals(new Inventory(99, "blue"), result.get(1, TimeUnit.SECONDS)); + } + + @Test + void concurrentWaitsAreCorrelated() throws Exception { + when(rpc.invoke(eq("session.send"), any(), eq(SendMessageResponse.class))).thenReturn( + CompletableFuture.completedFuture(new SendMessageResponse("first")), + CompletableFuture.completedFuture(new SendMessageResponse("second"))); + var first = session.sendAndWait("one", Inventory.class); + var second = session.sendAndWait("two", Inventory.class); + assistant("first", "{\"count\":42,\"color\":\"red\"}"); + assistant("second", "{\"count\":7,\"color\":\"blue\"}"); + emit("session.idle", Map.of()); + assertEquals(42, first.get(1, TimeUnit.SECONDS).count()); + assertEquals(7, second.get(1, TimeUnit.SECONDS).count()); + } + + @ParameterizedTest + @ValueSource(strings = {"null", "not JSON", "{\"count\":\"bad\"}", "{\"count\":42,\"color\":\"red\"} trailing"}) + void rejectsNullAndInvalidJson(String content) throws Exception { + var result = session.sendAndWait("inventory", Inventory.class); + assistant("user-1", content); + emit("session.idle", Map.of()); + assertThrows(ExecutionException.class, () -> result.get(1, TimeUnit.SECONDS)); + } + + @ParameterizedTest + @ValueSource(strings = {"missing", "aborted", "error", "tool"}) + void rejectsFailedOrIncompleteRuns(String kind) throws Exception { + var result = session.sendAndWait("inventory", Inventory.class); + emit("user.message", Map.of("messageId", "user-1", "content", "inventory")); + if (kind.equals("tool")) { + assistant("user-1", "{\"count\":42,\"color\":\"red\"}"); + emit("assistant.message", Map.of("messageId", "commentary", "originatingMessageId", "user-1", "content", + "working", "toolRequests", List.of(Map.of("toolCallId", "t", "name", "tool")))); + } + if (kind.equals("error")) { + emit("session.error", Map.of("errorType", "test", "message", "provider failed")); + } else { + emit("session.idle", kind.equals("aborted") ? Map.of("aborted", true) : Map.of()); + } + assertThrows(ExecutionException.class, () -> result.get(1, TimeUnit.SECONDS)); + } + + @Test + void timeoutCancellationAndCloseCompleteWaits() throws Exception { + var timed = session.sendAndWait(new MessageOptions().setPrompt("inventory"), Inventory.class, 10); + assertThrows(ExecutionException.class, () -> timed.get(1, TimeUnit.SECONDS)); + var cancelled = session.sendAndWait("inventory", Inventory.class); + assertTrue(cancelled.cancel(true)); + var closing = session.sendAndWait(new MessageOptions().setPrompt("inventory"), Inventory.class, 0); + session.close(); + assertThrows(ExecutionException.class, () -> closing.get(1, TimeUnit.SECONDS)); + } + + @Test + void admissionFailureIsNotHiddenByBufferedEvents() throws Exception { + when(rpc.invoke(eq("session.send"), any(), eq(SendMessageResponse.class))).thenAnswer(invocation -> { + assistant("user-1", "{\"count\":42,\"color\":\"red\"}"); + emit("session.idle", Map.of()); + return CompletableFuture.failedFuture(new IllegalArgumentException("invalid schema")); + }); + var result = session.sendAndWait("inventory", Inventory.class); + assertEquals("invalid schema", + assertThrows(ExecutionException.class, () -> result.get(1, TimeUnit.SECONDS)).getCause().getMessage()); + } + + @Test + void explicitSchemaIsForwardedUnchangedAndConflictingTypedOptionsAreRejected() throws Exception { + var schema = Map.of("type", "object", "description", "unchanged"); + var options = new MessageOptions().setPrompt("inventory").setResponseSchema(schema); + assertThrows(IllegalArgumentException.class, () -> session.sendAndWait(options, Inventory.class)); + assertThrows(IllegalArgumentException.class, () -> session + .sendAndWait(new MessageOptions().setPrompt("inventory").setMode("immediate"), Inventory.class)); + verifyNoInteractions(rpc); + var result = session.sendAndWait(options); + assistant("user-1", "{}"); + emit("session.idle", Map.of()); + assertEquals("{}", result.get(1, TimeUnit.SECONDS).getData().content()); + verify(rpc).invoke(eq("session.send"), + argThat(request -> request instanceof SendMessageRequest send + && schema.equals(((Map) send.getResponseFormat().get("jsonSchema")).get("schema"))), + eq(SendMessageResponse.class)); + } +} diff --git a/java/sdk/src/test/java/com/github/copilot/tool/CopilotToolProcessorTest.java b/java/sdk/src/test/java/com/github/copilot/tool/CopilotToolProcessorTest.java index e7012c644f..86f11e3f45 100644 --- a/java/sdk/src/test/java/com/github/copilot/tool/CopilotToolProcessorTest.java +++ b/java/sdk/src/test/java/com/github/copilot/tool/CopilotToolProcessorTest.java @@ -53,6 +53,36 @@ class CopilotToolProcessorTest { @TempDir java.nio.file.Path tempDir; + @Test + void generatesClosedResponseSchemaForNestedRecords() { + var result = compileWithProcessor(List.of(inMemorySource("test.Result", """ + package test; + import com.github.copilot.CopilotResponse; + import java.util.List; + @CopilotResponse + public record Result(List items) { + public record Item(int count, String color) {} + } + """))); + assertNoErrors(result); + var generated = result.getGeneratedSource("test.Result$$CopilotResponseMeta"); + assertNotNull(generated); + assertTrue(generated.contains("\"additionalProperties\", false")); + assertTrue(generated.contains("\"count\", Map.of(\"type\", \"integer\")")); + } + + @Test + void reportsRecursiveResponseTypesInsteadOfOverflowing() { + var result = compileWithProcessor(List.of(inMemorySource("test.Result", """ + package test; + import com.github.copilot.CopilotResponse; + @CopilotResponse + public record Result(Result next) {} + """))); + assertTrue(result.diagnostics.stream().anyMatch(d -> d.getKind() == Diagnostic.Kind.ERROR + && d.getMessage(null).contains("Recursive response types require an explicit JSON Schema"))); + } + // ── Test: Basic generation ────────────────────────────────────────────────── @Test @@ -1420,7 +1450,8 @@ private CompilationResult compileWithProcessor(List sources) { String classpath = resolveClasspath(); List options = new ArrayList<>(); options.add("-proc:full"); - options.addAll(List.of("-processor", "com.github.copilot.tool.CopilotToolProcessor")); + options.addAll(List.of("-processor", + "com.github.copilot.tool.CopilotToolProcessor,com.github.copilot.tool.CopilotResponseProcessor")); options.addAll(List.of("-classpath", classpath)); options.addAll(List.of("-d", tempDir.toString())); options.addAll(List.of("-s", tempDir.toString())); diff --git a/nodejs/README.md b/nodejs/README.md index a6e577be93..93b6160808 100644 --- a/nodejs/README.md +++ b/nodejs/README.md @@ -305,6 +305,87 @@ Send a message and wait until the session becomes idle. Returns the final assistant message event, or undefined if none was received. +##### Structured output (preview) + +Requires a runtime build with `responseFormat` and `originatingMessageId` support. +Pass a raw JSON Schema or a Zod schema as `responseSchema` to `send` or +`sendAndWait`. As with custom tool parameters, the SDK converts Zod schemas to +JSON Schema before sending them: + +```typescript +import { z } from "zod"; + +const answerSchema = z.object({ answer: z.number().int() }); +const message = await session.sendAndWait({ + prompt: "What is 19 + 23?", + responseSchema: answerSchema, +}); +console.log(message?.data.content); // JSON text +``` + +For a typed result, pass the Zod schema as the **second argument** instead: + +```typescript +const answer = await session.sendAndWait("What is 19 + 23?", answerSchema); +console.log(answer.answer); // number; TResult is inferred from answerSchema +``` + +`sendAndWait(options, schema, timeout?)` generates the JSON Schema from +the schema value, parses the final JSON, and validates it with the schema's +`parse` method. TypeScript cannot derive a runtime schema from an erased type +parameter alone. Invalid JSON, a schema mismatch, or a completed run without a +matching assistant message throws. Do not also set `options.responseSchema` when +using the typed overload. + +The schema belongs to the submitted run, including its tool-call iterations. +Internally generated stop-hook corrections retain the schema and originating +message ID, so the wait returns the corrected answer. Independent subsequent +sends do not inherit it. Ordinary immediate steering inherits the active schema +and originating message ID, even when it arrives too late for the current model +request and is promoted into a follow-up run. Specifying a schema with +`mode: "immediate"` is rejected, even while idle. +The generated `session.rpc.send` and `session.rpc.sendMessages` wrappers expose +the full `responseFormat` contract when you need to set its name, description, +or strict option rather than using the convenience defaults (`name: "response"`, +`strict: true`). +Each batch starts one run: the final returned message ID is its origin, preceding +messages are context, and an empty batch has no origin. An immediate batch +steers the active run instead and retains its origin. +The schema is not a persisted session default: autonomous resume-pending work +after a restart does not restore it. A terminal tool that clears context ends +the old run; its fresh seed does not inherit the schema or origin. Such a run +can finish without a structured result, in which case the typed wait throws. +After a successful terminal tool, the runtime disables tools while the model +produces the structured result. Stop-hook corrections remain supported. +Remote sessions and known HydraFusion routes reject response formats before +admission. Schemas larger than 32 MiB when JSON-encoded are also rejected before +admission, using the runtime's existing request-size ceiling. This does not +guarantee the schema plus conversation and tools fits the provider's budget. + +Structured waits select the last root-agent message whose `originatingMessageId` +matches the ID returned by their send, then return at a non-autopilot +`session.idle`. Other queued work can delay that idle, but cannot replace the +selected result. The existing unformatted overload retains its session-wide +behavior. `turnId` identifies an individual model/tool iteration, not the whole +run; telemetry interaction IDs are not unique run identifiers. + +For event-driven consumption with `send`, subscribe before sending and collect +root `assistant.message` events whose `data.originatingMessageId` matches the ID +returned by `send`; events may arrive before that acknowledgement. Wait for +`session.idle`, then parse the last matching message without tool requests. +An earlier response may be superseded by a stop-hook correction. Handle +`session.error` and aborted idle events rather than returning a partial result. + +Streaming still delivers ordinary text events, including intermediate messages +and tool calls. Only the final selected message is parsed by the typed overload; +not every event is necessarily a complete schema-conforming JSON document. +Provider errors, refusals, cancellation, truncation, session errors, and timeouts +can prevent a typed result. A timeout stops waiting, not the runtime's work. +Use a model and endpoint that support native structured output. An API-compatible +gateway may ignore format fields even when it accepts the request; for example, +the Claude Chat-completions compatibility route is not equivalent to Anthropic's +native `output_config.format` endpoint. + ##### `on(eventType: string, handler: TypedSessionEventHandler): () => void` Subscribe to a specific event type. The handler receives properly typed events. diff --git a/nodejs/src/client.ts b/nodejs/src/client.ts index 6e4b4fb5b4..8438033e18 100644 --- a/nodejs/src/client.ts +++ b/nodejs/src/client.ts @@ -49,6 +49,7 @@ import { createSessionFsAdapter, type SessionFsProvider } from "./sessionFsProvi import { createCopilotRequestAdapter } from "./copilotRequestHandler.js"; import type { CopilotRequestHandler } from "./copilotRequestHandler.js"; import { getTraceContext } from "./telemetry.js"; +import { toJsonSchema } from "./schema.js"; import { ToolSet } from "./toolSet.js"; import type { AutoModeSwitchRequest, @@ -87,7 +88,6 @@ import type { SessionMetadata, SystemMessageCustomizeConfig, TelemetryConfig, - Tool, TraceContextProvider, TypedSessionLifecycleHandler, } from "./types.js"; @@ -101,18 +101,6 @@ import type { FactoryHandle } from "./factory.js"; const MIN_PROTOCOL_VERSION = 3; const RUNTIME_SHUTDOWN_TIMEOUT_MS = 10_000; -/** - * Check if value is a Zod schema (has toJSONSchema method) - */ -function isZodSchema(value: unknown): value is { toJSONSchema(): Record } { - return ( - value != null && - typeof value === "object" && - "toJSONSchema" in value && - typeof (value as { toJSONSchema: unknown }).toJSONSchema === "function" - ); -} - async function withTimeout(promise: Promise, timeoutMs: number, message: string): Promise { let timeout: ReturnType | undefined; try { @@ -160,17 +148,6 @@ async function waitForChildExit(child: ChildProcess, timeoutMs: number): Promise }); } -/** - * Convert tool parameters to JSON schema format for sending to CLI - */ -function toJsonSchema(parameters: Tool["parameters"]): Record | undefined { - if (!parameters) return undefined; - if (isZodSchema(parameters)) { - return parameters.toJSONSchema(); - } - return parameters; -} - /** Implicit provider name for the singular, whole-session {@link ProviderConfig}. */ const DEFAULT_PROVIDER_NAME = "default"; diff --git a/nodejs/src/index.ts b/nodejs/src/index.ts index 8a5a730b5a..a9244e9e48 100644 --- a/nodejs/src/index.ts +++ b/nodejs/src/index.ts @@ -116,6 +116,7 @@ export type { DefaultAgentConfig, BearerTokenProvider, MessageOptions, + ResponseSchema, MessageSource, ManagedSettings, ManagedSettingsPermissions, diff --git a/nodejs/src/schema.ts b/nodejs/src/schema.ts new file mode 100644 index 0000000000..29d9426758 --- /dev/null +++ b/nodejs/src/schema.ts @@ -0,0 +1,24 @@ +/*--------------------------------------------------------------------------------------------- + * Copyright (c) Microsoft Corporation. All rights reserved. + *--------------------------------------------------------------------------------------------*/ + +import type { ResponseSchema, ZodSchema } from "./types.js"; + +export function isZodSchema(value: unknown): value is ZodSchema { + return ( + typeof value === "object" && + value !== null && + "toJSONSchema" in value && + typeof value.toJSONSchema === "function" + ); +} + +export function toJsonSchema( + schema: ZodSchema | Record | undefined +): Record | undefined { + return isZodSchema(schema) ? schema.toJSONSchema() : schema; +} + +export function isResponseSchema(value: unknown): value is ResponseSchema { + return isZodSchema(value) && "parse" in value && typeof value.parse === "function"; +} diff --git a/nodejs/src/session.ts b/nodejs/src/session.ts index 896eded0bd..4a909f64c8 100644 --- a/nodejs/src/session.ts +++ b/nodejs/src/session.ts @@ -23,6 +23,7 @@ import type { import { type Canvas, CanvasError } from "./canvas.js"; import type { OpenCanvasInstance } from "./generated/rpc.js"; import { getTraceContext } from "./telemetry.js"; +import { isResponseSchema, toJsonSchema } from "./schema.js"; import { isAttributedPermissionResult } from "./types.js"; import type { CommandHandler, @@ -39,6 +40,7 @@ import type { BearerTokenProvider, UiInputOptions, MessageOptions, + ResponseSchema, McpAuthHandler, McpAuthRequest, PermissionHandler, @@ -458,6 +460,7 @@ export class CopilotSession { private _capabilities: SessionCapabilities = {}; private openCanvasInstances: OpenCanvasInstance[] = []; private disconnected = false; + private readonly pendingStructuredWaits = new Set<(error: Error) => void>(); private disconnecting = false; private onDisconnected?: () => void; @@ -721,13 +724,14 @@ export class CopilotSession { } /** - * Sends a message to this session and waits for the response. + * Sends a message to this session and returns once it is admitted. * * The message is processed asynchronously. Subscribe to events via {@link on} * to receive streaming responses and other session events. * * @param options - The message options including the prompt and optional attachments - * @returns A promise that resolves with the message ID of the response + * @returns The submitted user message's ID, not an assistant response ID. + * When this send starts a run, root assistant messages carry it as originatingMessageId. * @throws Error if the session has been disconnected or the connection fails * * @example @@ -753,6 +757,18 @@ export class CopilotSession { mode: options.mode, agentMode: options.agentMode, requestHeaders: options.requestHeaders, + ...(options.responseSchema + ? { + responseFormat: { + type: "json_schema", + jsonSchema: { + name: "response", + strict: true, + schema: toJsonSchema(options.responseSchema), + }, + }, + } + : {}), }); return (response as { messageId: string }).messageId; @@ -766,6 +782,9 @@ export class CopilotSession { * assistant has finished processing the message. * * Events are still delivered to handlers registered via {@link on} while waiting. + * With a schema as the second argument, returns its parsed, validated result. + * Structured waits select only root-agent output originating from this send; + * other queued work may delay session.idle but cannot replace the result. * * @param options - The message options including the prompt and optional attachments * @param timeout - Timeout in milliseconds (default: 60000). Controls how long to wait; does not abort in-flight agent work. @@ -782,17 +801,57 @@ export class CopilotSession { * ``` */ async sendAndWait(prompt: string, timeout?: number): Promise; + async sendAndWait( + options: MessageOptions | string, + responseSchema: ResponseSchema, + timeout?: number + ): Promise; async sendAndWait( options: MessageOptions, timeout?: number ): Promise; async sendAndWait( optionsOrPrompt: MessageOptions | string, + schemaOrTimeout?: ResponseSchema | number, timeout?: number - ): Promise { + ): Promise { const options: MessageOptions = typeof optionsOrPrompt === "string" ? { prompt: optionsOrPrompt } : optionsOrPrompt; - const effectiveTimeout = timeout ?? 60_000; + const typedSchema = isResponseSchema(schemaOrTimeout) ? schemaOrTimeout : undefined; + if (schemaOrTimeout !== undefined && typeof schemaOrTimeout !== "number" && !typedSchema) { + throw new TypeError( + "The second argument must be a timeout or a schema with toJSONSchema() and parse(). " + + "Pass raw JSON Schema in options.responseSchema instead." + ); + } + const effectiveTimeout = + (typeof schemaOrTimeout === "number" ? schemaOrTimeout : timeout) ?? 60_000; + + if (typedSchema && options.responseSchema) { + throw new Error( + "Do not specify responseSchema in options when requesting a typed response." + ); + } + if (typedSchema && options.mode === "immediate") { + throw new Error( + "Structured output cannot be requested on an immediate steering message." + ); + } + if (typedSchema || options.responseSchema) { + const message = await this.sendAndWaitForStructuredMessage( + typedSchema ? { ...options, responseSchema: typedSchema } : options, + effectiveTimeout + ); + if (typedSchema) { + if (!message) { + throw new Error( + "The requested run completed without a structured assistant response." + ); + } + return typedSchema.parse(JSON.parse(message.data.content)); + } + return message; + } type SessionOutcome = { kind: "idle" } | { kind: "error"; error: Error }; let resolveOutcome: (outcome: SessionOutcome) => void; @@ -845,12 +904,114 @@ export class CopilotSession { } } + private async sendAndWaitForStructuredMessage( + options: MessageOptions, + timeout: number + ): Promise { + if (this.disconnected) { + throw new Error("Session is disconnected"); + } + type Outcome = + | { kind: "idle"; message: AssistantMessageEvent | undefined } + | { kind: "error"; error: Error }; + let resolveOutcome!: (outcome: Outcome) => void; + const outcomePromise = new Promise((resolve) => { + resolveOutcome = resolve; + }); + const fail = (error: Error) => resolveOutcome({ kind: "error", error }); + let messageId: string | undefined; + let consumed = false; + let lastMessage: AssistantMessageEvent | undefined; + const buffered: SessionEvent[] = []; + const observe = (event: SessionEvent) => { + if (event.agentId) return; + if (event.type === "user.message" && event.data.messageId === messageId) { + consumed = true; + } else if ( + event.type === "assistant.message" && + event.data.originatingMessageId === messageId + ) { + consumed = true; + lastMessage = event.data.toolRequests?.length ? undefined : event; + } else if ( + consumed && + event.type === "session.idle" && + event.data.mode !== "autopilot" + ) { + if (event.data.aborted) { + fail( + new Error( + "The requested run was aborted before a structured result was completed." + ) + ); + } else { + resolveOutcome({ kind: "idle", message: lastMessage }); + } + } else if (consumed && event.type === "session.error") { + const error = new Error(event.data.message); + error.stack = event.data.stack; + fail(error); + } + }; + const unsubscribe = this.on((event) => { + if ( + event.type !== "user.message" && + event.type !== "assistant.message" && + event.type !== "session.idle" && + event.type !== "session.error" + ) { + return; + } + if (messageId === undefined) { + buffered.push(event); + } else { + observe(event); + } + }); + this.pendingStructuredWaits.add(fail); + const timer = setTimeout( + () => fail(new Error(`Timeout after ${timeout}ms waiting for the structured response`)), + timeout + ); + try { + const sendOutcome = this.send(options).then( + (id) => { + if (!id) { + throw new Error( + "The runtime did not return a message ID for the structured send." + ); + } + messageId = id; + for (const event of buffered) observe(event); + buffered.length = 0; + return outcomePromise; + }, + (error: unknown): Outcome => ({ + kind: "error", + error: error instanceof Error ? error : new Error(String(error)), + }) + ); + const outcome = await Promise.race([sendOutcome, outcomePromise]); + if (outcome.kind === "error") throw outcome.error; + return outcome.message; + } finally { + clearTimeout(timer); + buffered.length = 0; + unsubscribe(); + this.pendingStructuredWaits.delete(fail); + } + } + /** @internal */ _markDisconnected(): void { if (this.disconnected) { return; } this.disconnected = true; + for (const fail of this.pendingStructuredWaits) { + fail(new Error("Session disconnected while waiting for a structured response")); + } + this.pendingStructuredWaits.clear(); for (const controller of this.pendingExternalTools.values()) { controller.abort(); } diff --git a/nodejs/src/types.ts b/nodejs/src/types.ts index 9c4258d9c9..de332d4a3a 100644 --- a/nodejs/src/types.ts +++ b/nodejs/src/types.ts @@ -710,6 +710,14 @@ export interface ZodSchema { toJSONSchema(): Record; } +/** + * A Zod-compatible output schema that both describes and parses a typed result. + * TypeScript types are erased at runtime, so typed output requires a schema value. + */ +export interface ResponseSchema extends ZodSchema { + parse(value: unknown): T; +} + /** * Tool definition. Parameters can be either: * - A Zod schema (provides type inference for handler) @@ -3395,6 +3403,20 @@ export interface MessageOptions { * If provided, this is shown in the timeline instead of `prompt`. */ displayPrompt?: string; + + /** + * JSON Schema or a Zod schema for this run's output, including requests after tool calls. + * Independent sends do not inherit it. Ordinary immediate steering retains the active + * schema and origin, even when promoted to a follow-up after the model request finishes. + * Specifying a schema with mode "immediate" is rejected, even while idle. + * This is not a persisted session default and does not survive a context reset. + * + * sendAndWait still returns an assistant message event. For a typed result, pass a + * Zod-compatible schema as sendAndWait's second argument instead. + * Streaming events remain text and may include intermediate messages. + * Use rpc.send's responseFormat for provider-specific name, description and strict options. + */ + responseSchema?: ZodSchema | Record; } /** diff --git a/nodejs/test/e2e/structured_output.e2e.test.ts b/nodejs/test/e2e/structured_output.e2e.test.ts new file mode 100644 index 0000000000..0eb5957a61 --- /dev/null +++ b/nodejs/test/e2e/structured_output.e2e.test.ts @@ -0,0 +1,502 @@ +/*--------------------------------------------------------------------------------------------- + * Copyright (c) Microsoft Corporation. All rights reserved. + *--------------------------------------------------------------------------------------------*/ + +import { describe, expect, expectTypeOf, it } from "vitest"; +import { z } from "zod"; +import { + approveAll, + defineTool, + type AssistantMessageEvent, + type CopilotSession, + type ProviderConfig, + type SessionEvent, +} from "../../src/index.js"; +import { createSdkTestContext, DEFAULT_GITHUB_TOKEN, isCI } from "./harness/sdkTestContext"; +import { waitForCondition } from "./harness/sdkTestHelper"; + +describe("Structured output", async () => { + const { copilotClient: client, openAiEndpoint } = await createSdkTestContext({ + copilotClientOptions: { + env: { COPILOT_CLI_ENABLED_FEATURE_FLAGS: "HYDRAFUSION,HYDRAFUSION_ROLLOUT" }, + }, + }); + const provider: ProviderConfig = { + type: "openai", + wireApi: "completions", + baseUrl: openAiEndpoint.url, + modelId: "gpt-4.1", + wireModel: "gpt-4.1", + apiKey: isCI ? DEFAULT_GITHUB_TOKEN : (process.env.GITHUB_TOKEN ?? DEFAULT_GITHUB_TOKEN), + headers: { + "Copilot-Integration-Id": "copilot-developer-cli", + "Copilot-Harness-Id": "copilot-sdk", + "X-GitHub-Api-Version": "2026-08-01", + }, + }; + + it("infers_typed_result_after_custom_tool", async () => { + let calls = 0; + const schema = z.object({ count: z.number().int(), color: z.string() }); + const session = await client.createSession({ + model: "gpt-4.1", + provider, + onPermissionRequest: approveAll, + availableTools: [], + streaming: true, + tools: [ + defineTool("get_inventory", { + description: "Get the current widget inventory.", + parameters: z.object({}), + handler: () => { + calls++; + return "The inventory contains 42 red widgets."; + }, + }), + ], + }); + const events: SessionEvent[] = []; + session.on((event) => events.push(event)); + const result = await session.sendAndWait( + "Call get_inventory, then report the widget count and color.", + schema + ); + expectTypeOf(result).toEqualTypeOf<{ count: number; color: string }>(); + expect(result).toEqual({ count: 42, color: "red" }); + expect(calls).toBeGreaterThan(0); + expect(events.some((event) => event.type === "assistant.message_delta")).toBe(true); + const ordinary = await session.sendAndWait( + "Now reply with exactly the plain text HELLO, not JSON." + ); + expect(ordinary?.data.content.trim()).toBe("HELLO"); + const exchanges = await openAiEndpoint.getExchanges(); + expect(exchanges.length).toBeGreaterThanOrEqual(3); + for (const exchange of exchanges.slice(0, -1)) { + expect(exchange.request).toHaveProperty( + "response_format.json_schema.schema", + schema.toJSONSchema() + ); + } + expect(exchanges.at(-1)!.request).not.toHaveProperty("response_format"); + }); + + it("typed_result_after_terminal_tool_and_steering", async () => { + const events: SessionEvent[] = []; + let calls = 0; + let session: CopilotSession; + session = await client.createSession({ + model: "gpt-4.1", + provider, + onPermissionRequest: approveAll, + availableTools: [], + streaming: true, + onEvent: (event) => events.push(event), + tools: [ + defineTool("lookup_number", { + description: "Return the number needed for the calculation.", + parameters: z.object({}), + skipPermission: true, + isTerminal: true, + handler: async () => { + calls++; + await session.send({ + prompt: "Continue with the original calculation. Do not call any more tools.", + mode: "immediate", + }); + return 58; + }, + }), + ], + }); + const schema = z.object({ answer: z.number().int(), contract: z.literal("typed_tool") }); + const result = await session.sendAndWait( + "Call lookup_number exactly once, then add 5 to the returned number. Do not guess its result.", + schema + ); + expectTypeOf(result).toEqualTypeOf<{ answer: number; contract: "typed_tool" }>(); + expect(result).toEqual({ answer: 63, contract: "typed_tool" }); + expect(calls).toBe(1); + expect(events.some((event) => event.type === "tool.execution_complete")).toBe(true); + expect(events.some((event) => event.type === "assistant.message_delta")).toBe(true); + const replies = events.filter( + (event) => event.type === "assistant.message" && !event.agentId + ); + expect(replies.some((event) => event.data.toolRequests?.length)).toBe(true); + expect(replies.at(-1)?.data.toolRequests ?? []).toEqual([]); + const exchanges = await openAiEndpoint.getExchanges(); + expect(exchanges.length).toBeGreaterThanOrEqual(2); + for (const exchange of exchanges.slice(1)) { + expect(exchange.request).toHaveProperty("tool_choice", "none"); + } + for (const exchange of exchanges) { + expect(exchange.request).toHaveProperty( + "response_format.json_schema.schema", + schema.toJSONSchema() + ); + } + }); + + it("typed_wait_returns_stop_hook_correction_after_terminal_tool", async () => { + let calls = 0; + let stops = 0; + const replies: AssistantMessageEvent[] = []; + const session = await client.createSession({ + model: "gpt-4.1", + provider, + onPermissionRequest: approveAll, + availableTools: [], + tools: [ + defineTool("lookup_number", { + description: "Return the number needed for the calculation.", + parameters: z.object({}), + skipPermission: true, + isTerminal: true, + handler: () => { + calls++; + return 58; + }, + }), + ], + onEvent: (event) => { + if (event.type === "assistant.message" && !event.agentId) replies.push(event); + }, + hooks: { + onAgentStop: () => + ++stops === 1 + ? { + decision: "block", + reason: "Correct the answer to 99, not 63. Do not use tools.", + } + : undefined, + }, + }); + const schema = z.object({ answer: z.number().int() }); + const result = await session.sendAndWait( + "Call lookup_number exactly once, then add 5 to the returned number. Do not guess its result.", + schema + ); + expect(result).toEqual({ answer: 99 }); + expect(calls).toBe(1); + expect(stops).toBe(2); + const answers = replies.filter((reply) => !reply.data.toolRequests?.length); + expect(answers.map((reply): unknown => JSON.parse(reply.data.content))).toEqual([ + { answer: 63 }, + { answer: 99 }, + ]); + expect(answers[0].data.originatingMessageId).toBeTruthy(); + expect(answers[1].data.originatingMessageId).toBe(answers[0].data.originatingMessageId); + const exchanges = await openAiEndpoint.getExchanges(); + expect(exchanges).toHaveLength(3); + expect(exchanges[1].request).toHaveProperty("tool_choice", "none"); + for (const exchange of exchanges) { + expect(exchange.request).toHaveProperty( + "response_format.json_schema.schema", + schema.toJSONSchema() + ); + } + }); + + it("rejects_unsupported_or_oversized_schemas_before_admission", async () => { + for (const model of ["gpt-4.1", "hydrafusion"]) { + const session = await client.createSession({ + model, + provider, + onPermissionRequest: approveAll, + availableTools: [], + }); + await expect( + session.sendAndWait( + { prompt: "Must not be admitted", mode: "immediate" }, + z.object({ answer: z.number().int() }) + ) + ).rejects.toThrow(/immediate/); + const schema = { + type: "object", + description: model === "gpt-4.1" ? "x".repeat(32 * 1024 * 1024) : "Small schema", + }; + const message = model === "gpt-4.1" ? /32 MiB/ : /HydraFusion/; + await expect( + session.sendAndWait({ + prompt: "Must not be admitted", + responseSchema: schema, + }) + ).rejects.toThrow(message); + await expect( + session.rpc.sendMessages({ + messages: [], + responseFormat: { + type: "json_schema", + jsonSchema: { name: "response", schema }, + }, + }) + ).rejects.toThrow(message); + expect((await session.rpc.queue.pendingItems()).items).toEqual([]); + expect( + (await session.getEvents()).filter( + (event) => event.type === "user.message" || event.type === "session.error" + ) + ).toEqual([]); + } + expect(await openAiEndpoint.getExchanges()).toEqual([]); + }); + + it("send_selects_correlated_response_after_idle", async () => { + let releaseHook!: () => void; + let hookEntered = false; + const hookReleased = new Promise((resolve) => { + releaseHook = resolve; + }); + const session = await client.createSession({ + model: "gpt-4.1", + provider, + onPermissionRequest: approveAll, + availableTools: [], + tools: [ + defineTool("read_inventory", { + description: "Read the current widget count and color.", + parameters: z.object({}), + skipPermission: true, + handler: () => "The inventory contains 42 red widgets.", + }), + ], + hooks: { + onAgentStop: async () => { + hookEntered = true; + await hookReleased; + }, + }, + }); + const replies: AssistantMessageEvent[] = []; + const errors: string[] = []; + let idle = false; + const unsubscribe = session.on((event) => { + if (event.agentId) return; + if (event.type === "assistant.message") { + replies.push(event); + } else if (event.type === "session.error") { + errors.push(event.data.message); + } else if (event.type === "session.idle") { + idle = true; + } + }); + const schema = z.object({ count: z.number().int(), color: z.string() }); + try { + const messageId = await session.send({ + prompt: "Call read_inventory once, then report the current widget count and color.", + responseSchema: schema, + }); + await waitForCondition(() => hookEntered || errors.length > 0, { + timeoutMessage: "Stop hook did not start", + }); + expect(errors).toEqual([]); + expect(idle).toBe(false); + releaseHook(); + await waitForCondition(() => idle || errors.length > 0, { + timeoutMessage: "Session did not become idle", + }); + expect(errors).toEqual([]); + const reply = replies.findLast( + (event) => event.data.originatingMessageId === messageId + ); + expect(reply).toBeDefined(); + if (!reply) throw new Error("No correlated assistant response"); + expect(reply.data.originatingMessageId).toBe(messageId); + expect(schema.parse(JSON.parse(reply.data.content))).toEqual({ + count: 42, + color: "red", + }); + expect(replies.some((event) => event.data.toolRequests?.length)).toBe(true); + expect(reply.data.toolRequests ?? []).toEqual([]); + expect(replies.at(-1)).toBe(reply); + } finally { + releaseHook(); + unsubscribe(); + } + }, 60_000); + + it("typed_wait_returns_stop_hook_correction", async () => { + let stops = 0; + const replies: AssistantMessageEvent[] = []; + const schema = z.object({ answer: z.number().int() }); + const session = await client.createSession({ + model: "gpt-4.1", + provider, + onPermissionRequest: approveAll, + availableTools: [], + onEvent: (event) => { + if (event.type === "assistant.message" && !event.agentId) replies.push(event); + }, + hooks: { + onAgentStop: () => + ++stops === 1 + ? { + decision: "block", + reason: "Correct the answer to 99, not 42. Do not use tools.", + } + : undefined, + }, + }); + const result = await session.sendAndWait("What is 19 + 23? Do not use tools.", schema); + expect(result).toEqual({ answer: 99 }); + expect(stops).toBe(2); + expect(replies.map((reply): unknown => JSON.parse(reply.data.content))).toEqual([ + { answer: 42 }, + { answer: 99 }, + ]); + expect(replies[0].data.originatingMessageId).toBeTruthy(); + expect(replies[1].data.originatingMessageId).toBe(replies[0].data.originatingMessageId); + const exchanges = await openAiEndpoint.getExchanges(); + expect(exchanges).toHaveLength(2); + for (const exchange of exchanges) { + expect(exchange.request).toHaveProperty( + "response_format.json_schema.schema", + schema.toJSONSchema() + ); + } + }); + + it("typed_wait_returns_late_steering_response", async () => { + let stops = 0; + let steeringId: string | undefined; + let session: CopilotSession; + const replies: AssistantMessageEvent[] = []; + const schema = z.object({ answer: z.number().int() }); + session = await client.createSession({ + model: "gpt-4.1", + provider, + onPermissionRequest: approveAll, + availableTools: [], + onEvent: (event) => { + if (event.type === "assistant.message" && !event.agentId) replies.push(event); + }, + hooks: { + onAgentStop: async () => { + if (++stops === 1) { + // The final model request has finished, but this run still admits steering. + steeringId = await session.send({ + prompt: "Change the answer to 99. Do not use tools.", + mode: "immediate", + }); + } + }, + }, + }); + const result = await session.sendAndWait("What is 19 + 23? Do not use tools.", schema); + expect(result).toEqual({ answer: 99 }); + expect(stops).toBe(2); + expect(replies.map((reply): unknown => JSON.parse(reply.data.content))).toEqual([ + { answer: 42 }, + { answer: 99 }, + ]); + expect(steeringId).toBeTruthy(); + expect(replies[0].data.originatingMessageId).toBeTruthy(); + expect(replies[0].data.originatingMessageId).not.toBe(steeringId); + expect(replies[1].data.originatingMessageId).toBe(replies[0].data.originatingMessageId); + const exchanges = await openAiEndpoint.getExchanges(); + expect(exchanges).toHaveLength(2); + for (const exchange of exchanges) { + expect(exchange.request).toHaveProperty( + "response_format.json_schema.schema", + schema.toJSONSchema() + ); + } + }); + + it("concurrent_typed_sends_return_their_own_results", async () => { + let markToolEntered!: () => void; + let releaseTool!: () => void; + const toolEntered = new Promise((resolve) => { + markToolEntered = resolve; + }); + const toolReleased = new Promise((resolve) => { + releaseTool = resolve; + }); + const session = await client.createSession({ + model: "gpt-4.1", + provider, + onPermissionRequest: approveAll, + availableTools: [], + tools: [ + defineTool("first_number", { + description: "Get the number for the first question.", + parameters: z.object({}), + skipPermission: true, + handler: async () => { + markToolEntered(); + await toolReleased; + return 42; + }, + }), + ], + }); + const first = session.sendAndWait( + "Call first_number exactly once and report its returned number.", + z.object({ first: z.number().int() }) + ); + try { + await Promise.race([ + toolEntered, + first.then(() => { + throw new Error("First run completed without calling first_number"); + }), + ]); + const secondPrompt = "What is 30 + 7? Do not use tools."; + const second = session.sendAndWait( + secondPrompt, + z.object({ second: z.number().int() }) + ); + const results = Promise.all([first, second]); + await Promise.race([ + waitForCondition( + async () => + (await session.rpc.queue.pendingItems()).items.some((item) => + item.displayText.includes(secondPrompt) + ), + { timeoutMessage: "Second structured send was not queued behind the tool call" } + ), + results.then(() => { + throw new Error("Runs completed before the tool was released"); + }), + ]); + releaseTool(); + const [firstResult, secondResult] = await results; + expect(firstResult).toEqual({ first: 42 }); + expect(secondResult).toEqual({ second: 37 }); + } finally { + releaseTool(); + } + }); + + it("sends_explicit_schema_for_message_and_batch", async () => { + const session = await client.createSession({ + model: "gpt-4.1", + provider, + onPermissionRequest: approveAll, + availableTools: [], + }); + const schema = z.object({ count: z.number().int(), color: z.string() }); + const events: SessionEvent[] = []; + session.on((event) => events.push(event)); + const response = await session.rpc.sendMessages({ + messages: [ + { prompt: "There are 42 red widgets in stock." }, + { prompt: "Report the widget count and color." }, + ], + responseFormat: { + type: "json_schema", + jsonSchema: { name: "inventory", strict: true, schema: schema.toJSONSchema() }, + }, + wait: true, + }); + const final = events.findLast((event) => event.type === "assistant.message"); + expect(final?.type).toBe("assistant.message"); + if (final?.type !== "assistant.message") throw new Error("No assistant response"); + expect(schema.parse(JSON.parse(final.data.content))).toEqual({ count: 42, color: "red" }); + expect(final.data.originatingMessageId).toBe(response.messageIds.at(-1)); + const raw = await session.sendAndWait({ + prompt: "The inventory now has 21 blue widgets. Report the new count and color.", + responseSchema: schema.toJSONSchema(), + }); + expect(schema.parse(JSON.parse(raw!.data.content))).toEqual({ count: 21, color: "blue" }); + }); +}); diff --git a/nodejs/test/structured-output.test.ts b/nodejs/test/structured-output.test.ts new file mode 100644 index 0000000000..8f029d0ac9 --- /dev/null +++ b/nodejs/test/structured-output.test.ts @@ -0,0 +1,307 @@ +/*--------------------------------------------------------------------------------------------- + * Copyright (c) Microsoft Corporation. All rights reserved. + *--------------------------------------------------------------------------------------------*/ + +import { afterEach, describe, expect, expectTypeOf, it, vi } from "vitest"; +import type { MessageConnection } from "vscode-jsonrpc/node.js"; +import { z } from "zod"; +import { CopilotSession } from "../src/session.js"; +import type { SessionEvent } from "../src/generated/session-events.js"; +import type { MessageOptions } from "../src/types.js"; + +const answer = z.object({ answer: z.number().int() }); + +function event(type: SessionEvent["type"], data: unknown, agentId?: string): SessionEvent { + return { + type, + data, + agentId, + id: crypto.randomUUID(), + timestamp: new Date().toISOString(), + parentId: null, + } as SessionEvent; +} + +function user(messageId: string): SessionEvent { + return event("user.message", { messageId, content: "question", turnId: "0" }); +} + +function assistant(originatingMessageId: string, content: string, agentId?: string): SessionEvent { + return event( + "assistant.message", + { messageId: crypto.randomUUID(), originatingMessageId, content, turnId: "1" }, + agentId + ); +} + +function controlledSession() { + const sends: Array<{ + params: Record; + resolve: (value: { messageId: string }) => void; + reject: (error: Error) => void; + }> = []; + const sendRequest = vi.fn((_method: string, params: Record) => { + return new Promise<{ messageId: string }>((resolve, reject) => { + sends.push({ params, resolve, reject }); + }); + }); + const session = new CopilotSession("session", { sendRequest } as unknown as MessageConnection); + return { session, sends, sendRequest }; +} + +async function sent(sends: unknown[], count = 1) { + await vi.waitFor(() => expect(sends).toHaveLength(count)); +} + +describe("structured output", () => { + afterEach(() => vi.useRealTimers()); + + it("infers TResult from a Zod schema and forwards its JSON Schema", async () => { + const { session, sends } = controlledSession(); + const pending = session.sendAndWait("What is 19 + 23?", answer); + expectTypeOf(pending).toEqualTypeOf>(); + await sent(sends); + expect(sends[0].params.responseFormat).toEqual({ + type: "json_schema", + jsonSchema: { name: "response", strict: true, schema: answer.toJSONSchema() }, + }); + session._dispatchEvent(user("one")); + session._dispatchEvent(assistant("one", '{"answer":42}')); + session._dispatchEvent(event("session.idle", {})); + sends[0].resolve({ messageId: "one" }); + await expect(pending).resolves.toEqual({ answer: 42 }); + }); + + it("keeps raw schema sends and options-based Zod sends event-shaped", async () => { + for (const schema of [answer.toJSONSchema(), answer]) { + const { session, sends } = controlledSession(); + const options: MessageOptions = { prompt: "question", responseSchema: schema }; + const pending = session.sendAndWait(options); + await sent(sends); + const final = assistant("one", '{"answer":42}'); + sends[0].resolve({ messageId: "one" }); + session._dispatchEvent(user("one")); + session._dispatchEvent(final); + session._dispatchEvent(event("session.idle", {})); + await expect(pending).resolves.toEqual(final); + } + }); + + it("isolates queued concurrent sends and excludes subagent messages", async () => { + const { session, sends } = controlledSession(); + const first = session.sendAndWait("first", answer); + const second = session.sendAndWait("second", answer); + await sent(sends, 2); + session._dispatchEvent(event("session.idle", {})); + session._dispatchEvent(user("one")); + session._dispatchEvent(assistant("one", "intermediate tool-call text")); + session._dispatchEvent(assistant("one", '{"answer":42}')); + session._dispatchEvent(user("two")); + session._dispatchEvent(assistant("two", '{"answer":37}')); + session._dispatchEvent(assistant("one", '{"answer":999}', "subagent")); + session._dispatchEvent(assistant("unrelated", '{"answer":123}')); + session._dispatchEvent(event("session.idle", {})); + sends[1].resolve({ messageId: "two" }); + sends[0].resolve({ messageId: "one" }); + await expect(first).resolves.toEqual({ answer: 42 }); + await expect(second).resolves.toEqual({ answer: 37 }); + }); + + it("freezes the final message at idle even when more events precede send acknowledgement", async () => { + const { session, sends } = controlledSession(); + const pending = session.sendAndWait("question", answer); + await sent(sends); + session._dispatchEvent(user("one")); + session._dispatchEvent(assistant("one", '{"answer":42}')); + session._dispatchEvent(event("session.idle", {})); + session._dispatchEvent(assistant("one", '{"answer":999}')); + sends[0].resolve({ messageId: "one" }); + await expect(pending).resolves.toEqual({ answer: 42 }); + }); + + it("ignores autopilot idle boundaries until a final idle", async () => { + const { session, sends } = controlledSession(); + const pending = session.sendAndWait("question", answer); + await sent(sends); + session._dispatchEvent(user("one")); + session._dispatchEvent(event("session.idle", { mode: "autopilot" })); + session._dispatchEvent(assistant("one", '{"answer":42}')); + session._dispatchEvent(event("session.idle", {})); + sends[0].resolve({ messageId: "one" }); + await expect(pending).resolves.toEqual({ answer: 42 }); + }); + + it.each(["refusal", '{"answer":"not a number"}', "null"])( + "rejects invalid final output: %s", + async (content) => { + const { session, sends } = controlledSession(); + const pending = session.sendAndWait("question", answer); + const assertion = expect(pending).rejects.toThrow(); + await sent(sends); + session._dispatchEvent(user("one")); + session._dispatchEvent(assistant("one", content)); + session._dispatchEvent(event("session.idle", {})); + sends[0].resolve({ messageId: "one" }); + await assertion; + } + ); + + it("rejects missing or uncorrelated output rather than borrowing another message", async () => { + const { session, sends } = controlledSession(); + const pending = session.sendAndWait("question", answer); + const assertion = expect(pending).rejects.toThrow( + "without a structured assistant response" + ); + await sent(sends); + session._dispatchEvent(user("one")); + session._dispatchEvent(assistant("two", '{"answer":42}')); + session._dispatchEvent(event("session.idle", {})); + sends[0].resolve({ messageId: "one" }); + await assertion; + }); + + it("rejects conflicting explicit and inferred schemas before sending", async () => { + const { session, sendRequest } = controlledSession(); + await expect( + session.sendAndWait({ prompt: "question", responseSchema: answer }, answer) + ).rejects.toThrow("Do not specify responseSchema"); + expect(sendRequest).not.toHaveBeenCalled(); + }); + + it("rejects typed immediate steering before sending", async () => { + const { session, sendRequest } = controlledSession(); + await expect( + session.sendAndWait({ prompt: "question", mode: "immediate" }, answer) + ).rejects.toThrow( + "Structured output cannot be requested on an immediate steering message." + ); + expect(sendRequest).not.toHaveBeenCalled(); + }); + + it.each([{ type: "object" }, { toJSONSchema: () => ({ type: "object" }) }, null])( + "rejects an invalid second argument instead of sending an unformatted request: %j", + async (schema) => { + const { session, sendRequest } = controlledSession(); + await expect( + // @ts-expect-error Exercise malformed arguments from JavaScript callers. + session.sendAndWait("question", schema) + ).rejects.toThrow("Pass raw JSON Schema in options.responseSchema instead."); + expect(sendRequest).not.toHaveBeenCalled(); + } + ); + + it("does not return a partial result after abort", async () => { + const { session, sends } = controlledSession(); + const pending = session.sendAndWait("question", answer); + const assertion = expect(pending).rejects.toThrow("aborted"); + await sent(sends); + session._dispatchEvent(user("one")); + session._dispatchEvent(assistant("one", '{"answer":42}')); + session._dispatchEvent(event("session.idle", { aborted: true })); + sends[0].resolve({ messageId: "one" }); + await assertion; + }); + + it("does not parse a tool-call message as the final result", async () => { + const { session, sends } = controlledSession(); + const pending = session.sendAndWait("question", answer); + const assertion = expect(pending).rejects.toThrow( + "without a structured assistant response" + ); + await sent(sends); + session._dispatchEvent(user("one")); + session._dispatchEvent( + event("assistant.message", { + messageId: "assistant-one", + originatingMessageId: "one", + content: '{"answer":42}', + toolRequests: [{ toolCallId: "tool-one", name: "lookup" }], + }) + ); + session._dispatchEvent(event("session.idle", {})); + sends[0].resolve({ messageId: "one" }); + await assertion; + }); + + it("propagates send and model failures", async () => { + const first = controlledSession(); + const sendFailure = first.session.sendAndWait("question", answer); + const sendAssertion = expect(sendFailure).rejects.toThrow("admission failed"); + await sent(first.sends); + first.sends[0].reject(new Error("admission failed")); + await sendAssertion; + + const second = controlledSession(); + const modelFailure = second.session.sendAndWait("question", answer); + const modelAssertion = expect(modelFailure).rejects.toThrow("provider rejected"); + await sent(second.sends); + second.session._dispatchEvent(user("one")); + second.session._dispatchEvent( + event("session.error", { message: "provider rejected", errorType: "query" }) + ); + second.sends[0].resolve({ messageId: "one" }); + await modelAssertion; + }); + + it("times out even while send acknowledgement is pending", async () => { + vi.useFakeTimers(); + const { session } = controlledSession(); + const pending = session.sendAndWait("question", answer, 100); + const assertion = expect(pending).rejects.toThrow("Timeout after 100ms"); + await vi.advanceTimersByTimeAsync(100); + await assertion; + }); + + it("does not treat an assistant response as successful completion", async () => { + const { session, sends } = controlledSession(); + const pending = session.sendAndWait("question", answer); + const assertion = expect(pending).rejects.toThrow("post-response failure"); + await sent(sends); + session._dispatchEvent(user("one")); + session._dispatchEvent( + event("assistant.message", { + messageId: "final-reply", + originatingMessageId: "one", + content: '{"answer":42}', + }) + ); + session._dispatchEvent( + event("session.error", { + errorType: "query", + message: "post-response failure", + }) + ); + session._dispatchEvent(event("session.idle", {})); + sends[0].resolve({ messageId: "one" }); + await assertion; + }); + + it("waits for idle and returns a correlated hook correction instead of the original answer", async () => { + const { session, sends } = controlledSession(); + const pending = session.sendAndWait("question", answer); + const completed = vi.fn(); + void pending.then(completed); + await sent(sends); + sends[0].resolve({ messageId: "one" }); + session._dispatchEvent(user("one")); + session._dispatchEvent(assistant("one", '{"answer":42}')); + await Promise.resolve(); + expect(completed).not.toHaveBeenCalled(); + session._dispatchEvent(user("hook-correction")); + session._dispatchEvent(assistant("one", '{"answer":99}')); + session._dispatchEvent(assistant("unrelated", '{"answer":123}')); + await Promise.resolve(); + expect(completed).not.toHaveBeenCalled(); + session._dispatchEvent(event("session.idle", {})); + await expect(pending).resolves.toEqual({ answer: 99 }); + }); + + it("rejects promptly when the session disconnects", async () => { + const { session, sends } = controlledSession(); + const pending = session.sendAndWait("question", answer); + const assertion = expect(pending).rejects.toThrow("Session disconnected"); + await sent(sends); + session._markDisconnected(); + await assertion; + }); +}); diff --git a/python/README.md b/python/README.md index 71a6e8ae78..8debd30676 100644 --- a/python/README.md +++ b/python/README.md @@ -542,6 +542,50 @@ Supported image formats include JPG, PNG, GIF, and other common image types. The await session.send("What does the most recent jpg in this directory portray?") ``` +## Structured output (experimental) + +Use a Pydantic model, just like custom-tool parameter schemas: + +```python +from pydantic import BaseModel, ConfigDict + + +class Inventory(BaseModel): + model_config = ConfigDict(extra="forbid") + count: int + color: str + + +inventory = await session.send_and_wait_typed( + "Call get_inventory, then report the widget count and color.", + Inventory, +) +print(inventory.count, inventory.color) +``` + +The helper derives a JSON Schema using `model_json_schema()` and validates the +final JSON with `model_validate_json(by_alias=True, by_name=False)` so validation +uses the schema's alias names, including in nested models, regardless of +model-level alias settings. This requires Pydantic 2.11 or newer. +For explicit schemas, use +`send(prompt, response_schema=schema)` or `send_and_wait(prompt, +response_schema=schema)`; the latter returns the ordinary message event. +`response_schema` also accepts a Pydantic model class without parsing the result. + +The schema applies to one run, including tools, steering, and stop-hook +corrections, not subsequent independent sends or subagents. Streaming remains +text. Structured waits select the last root message without tool requests whose +`originating_message_id` matches the admitted message, at non-autopilot idle. +Concurrent structured waits keep their own results; later queued work can delay +idle. Aborts, session errors after the run starts, or missing final output fail +the wait. Timeout/cancellation only stops waiting, not agent work. Immediate +steering cannot specify a schema. + +Provider schema restrictions still apply: use closed objects (as above) and +required fields for strict OpenAI output. Schemas are not rewritten; unsupported +models/schemas produce errors. Low-level `session.rpc.send` and +`session.rpc.send_messages` expose the full `ResponseFormat` options. + ## Streaming Enable streaming to receive assistant response chunks as they're generated: diff --git a/python/copilot/session.py b/python/copilot/session.py index 4149aa5893..4ce1862f48 100644 --- a/python/copilot/session.py +++ b/python/copilot/session.py @@ -21,7 +21,9 @@ from datetime import UTC, datetime from enum import Enum from types import TracebackType -from typing import TYPE_CHECKING, Any, Literal, NotRequired, Required, TypedDict, cast +from typing import TYPE_CHECKING, Any, Literal, NotRequired, Required, TypedDict, TypeVar, cast + +from pydantic import BaseModel from ._diagnostics import log_timing from ._jsonrpc import JsonRpcError, ProcessExitedError @@ -85,6 +87,7 @@ SessionEvent, SessionIdleData, SessionMode, + UserMessageData, session_event_from_dict, ) from .generated.session_events import ( @@ -99,6 +102,7 @@ ) logger = logging.getLogger(__name__) +TResponse = TypeVar("TResponse", bound=BaseModel) # Fixed name of the runtime's built-in tool-search tool. A client can replace # its behavior by registering a tool with this exact name and @@ -1655,6 +1659,7 @@ def __init__( self._open_canvases_lock = threading.Lock() self._rpc: SessionRpc | None = None self._destroyed = False + self._structured_waits: set[asyncio.Future[SessionEvent]] = set() self._disconnect_lock = asyncio.Lock() self._on_disconnect = on_disconnect @@ -1678,9 +1683,17 @@ def _cancel_pending_external_tools(self) -> None: def _mark_disconnected(self) -> None: self._destroyed = True + self._fail_structured_waits() self._cancel_pending_external_tools() self._run_disconnect_callback() + def _fail_structured_waits(self) -> None: + for future in tuple(self._structured_waits): + if not future.done(): + future.set_exception( + RuntimeError("Session closed before structured output completed") + ) + @property def rpc(self) -> SessionRpc: """Typed session-scoped RPC methods.""" @@ -1734,6 +1747,7 @@ async def send( agent_mode: Literal["interactive", "plan", "autopilot", "shell"] | None = None, request_headers: dict[str, str] | None = None, display_prompt: str | None = None, + response_schema: dict[str, Any] | type[BaseModel] | None = None, ) -> str: """ Send a message to this session. @@ -1755,6 +1769,8 @@ async def send( request_headers: Optional per-turn HTTP headers for outbound model requests. display_prompt: If provided, this is shown in the timeline instead of ``prompt``. + response_schema: JSON Schema or a Pydantic model for this run. Independent + sends do not inherit it. Immediate steering cannot specify a schema. Returns: The message ID assigned by the server, which can be used to correlate events. @@ -1786,6 +1802,16 @@ async def send( params["requestHeaders"] = request_headers if display_prompt is not None: params["displayPrompt"] = display_prompt + if response_schema is not None: + schema = ( + response_schema.model_json_schema() + if isinstance(response_schema, type) and issubclass(response_schema, BaseModel) + else response_schema + ) + params["responseFormat"] = { + "type": "json_schema", + "jsonSchema": {"name": "response", "strict": True, "schema": schema}, + } params.update(get_trace_context()) rpc_start = time.perf_counter() @@ -1811,6 +1837,7 @@ async def send_and_wait( agent_mode: Literal["interactive", "plan", "autopilot", "shell"] | None = None, request_headers: dict[str, str] | None = None, display_prompt: str | None = None, + response_schema: dict[str, Any] | type[BaseModel] | None = None, timeout: float = 60.0, ) -> SessionEvent | None: """ @@ -1837,6 +1864,8 @@ async def send_and_wait( ``prompt``. timeout: Timeout in seconds (default: 60). Controls how long to wait; does not abort in-flight agent work. + response_schema: A per-run schema. Waits for the last correlated root + assistant message without tool requests at non-autopilot idle. Returns: The final assistant message event, or None if none was received. @@ -1853,6 +1882,20 @@ async def send_and_wait( ... case AssistantMessageData() as data: ... print(data.content) """ + if response_schema is not None: + return await self._wait_for_structured_message( + lambda: self.send( + prompt, + attachments=attachments, + source=source, + mode=mode, + agent_mode=agent_mode, + request_headers=request_headers, + display_prompt=display_prompt, + response_schema=response_schema, + ), + timeout, + ) total_start = time.perf_counter() idle_event = asyncio.Event() error_event: Exception | None = None @@ -1931,6 +1974,125 @@ def handler(event: SessionEventTypeAlias) -> None: finally: unsubscribe() + async def send_and_wait_typed( + self, + prompt: str, + response_type: type[TResponse], + *, + attachments: list[Attachment] | None = None, + source: MessageSource | None = None, + mode: Literal["enqueue", "immediate"] | None = None, + agent_mode: Literal["interactive", "plan", "autopilot", "shell"] | None = None, + request_headers: dict[str, str] | None = None, + display_prompt: str | None = None, + timeout: float = 60.0, + ) -> TResponse: + """Infer a schema using Pydantic, then validate and return the final result. + + Uses the same model schema generation as custom tools. Provider schema + restrictions still apply (for example, configure ``extra="forbid"`` for + closed objects). Validation uses the schema's alias names, including for + nested models, regardless of model-level alias validation settings. + Streaming events remain text. Timeout or cancellation + only stops waiting, not the agent. Errors remain session-scoped. + """ + if not isinstance(response_type, type) or not issubclass(response_type, BaseModel): + raise TypeError("response_type must be a Pydantic BaseModel subclass") + if mode == "immediate": + raise ValueError( + "Structured output cannot be requested on an immediate steering message" + ) + response = await self.send_and_wait( + prompt, + attachments=attachments, + source=source, + mode=mode, + agent_mode=agent_mode, + request_headers=request_headers, + display_prompt=display_prompt, + response_schema=response_type, + timeout=timeout, + ) + assert response is not None and isinstance(response.data, AssistantMessageData) + return response_type.model_validate_json( + response.data.content, by_alias=True, by_name=False + ) + + async def _wait_for_structured_message( + self, send: Callable[[], Awaitable[str]], timeout: float + ) -> SessionEvent: + if self._destroyed: + raise RuntimeError("Session is disconnected") + completion: asyncio.Future[SessionEvent] = asyncio.get_running_loop().create_future() + self._structured_waits.add(completion) + pending: list[SessionEvent] = [] + message_id: str | None = None + started = False + final_message: SessionEvent | None = None + + def process(event: SessionEvent) -> None: + nonlocal started, final_message + if completion.done() or event.agent_id: + return + match event.data: + case UserMessageData() as data if data.message_id == message_id: + started = True + case AssistantMessageData() as data if data.originating_message_id == message_id: + started = True + final_message = None if data.tool_requests else event + case SessionIdleData() as data if started and data.mode != SessionMode.AUTOPILOT: + if data.aborted: + completion.set_exception( + RuntimeError("Session aborted before structured output completed") + ) + elif ( + final_message is None + or not cast(AssistantMessageData, final_message.data).content.strip() + ): + completion.set_exception( + RuntimeError("Run completed without a structured assistant response") + ) + else: + completion.set_result(final_message) + case SessionErrorData() as data if started: + completion.set_exception(RuntimeError(f"Session error: {data.message}")) + + def handler(event: SessionEvent) -> None: + if not isinstance( + event.data, + (UserMessageData, AssistantMessageData, SessionIdleData, SessionErrorData), + ): + return + if message_id is None: + pending.append(event) + else: + process(event) + + unsubscribe = self.on(handler) + admission: asyncio.Future[str] | None = None + try: + async with asyncio.timeout(timeout): + admission = asyncio.ensure_future(send()) + await asyncio.wait((admission, completion), return_when=asyncio.FIRST_COMPLETED) + if completion.done(): + return completion.result() + message_id = await admission + for event in pending: + process(event) + pending.clear() + return await completion + finally: + unsubscribe() + self._structured_waits.discard(completion) + if admission is not None: + admission.cancel() + await asyncio.gather(admission, return_exceptions=True) + # Retrieve an error if admission failed while disconnection also completed the future. + if completion.done() and not completion.cancelled(): + completion.exception() + else: + completion.cancel() + def on(self, handler: Callable[[SessionEvent], None]) -> Callable[[], None]: """ Subscribe to events from this session. @@ -3102,6 +3264,7 @@ async def disconnect(self) -> None: self._run_disconnect_callback() with self._event_handlers_lock: self._destroyed = True + self._fail_structured_waits() self._event_handlers.clear() with self._tool_handlers_lock: self._tool_handlers.clear() diff --git a/python/e2e/test_structured_output_e2e.py b/python/e2e/test_structured_output_e2e.py new file mode 100644 index 0000000000..e0bf27a2d9 --- /dev/null +++ b/python/e2e/test_structured_output_e2e.py @@ -0,0 +1,449 @@ +"""Structured-output tests against the released runtime and shared provider captures.""" + +import asyncio +import json + +import pytest +from pydantic import BaseModel, ConfigDict, Field + +from copilot import define_tool +from copilot.rpc import ( + JSONSchemaResponseFormat, + ResponseFormat, + ResponseFormatType, + SendMessageItem, + SendMessagesRequest, +) +from copilot.session import PermissionHandler +from copilot.session_events import ( + AssistantMessageData, + AssistantMessageDeltaData, + SessionErrorData, + SessionIdleData, + UserMessageData, +) + +pytestmark = pytest.mark.asyncio(loop_scope="module") + + +class Inventory(BaseModel): + model_config = ConfigDict(extra="forbid") + count: int + color: str + + +class Answer(BaseModel): + model_config = ConfigDict(extra="forbid") + answer: int + + +class ToolAnswer(BaseModel): + model_config = ConfigDict(extra="forbid") + answer: int + contract: str + + +class First(BaseModel): + model_config = ConfigDict(extra="forbid") + first: int + + +class Second(BaseModel): + model_config = ConfigDict(extra="forbid") + second: int + + +def config(ctx): + return { + "model": "gpt-4.1", + "available_tools": [], + "on_permission_request": PermissionHandler.approve_all, + "provider": { + "type": "openai", + "wire_api": "completions", + "base_url": ctx.proxy_url, + "model_id": "gpt-4.1", + "wire_model": "gpt-4.1", + "api_key": "fake-token-for-e2e-tests", + "headers": { + "Copilot-Integration-Id": "copilot-developer-cli", + "Copilot-Harness-Id": "copilot-sdk", + "X-GitHub-Api-Version": "2026-08-01", + }, + }, + } + + +async def test_infers_typed_result_after_custom_tool(ctx): + class AliasedInventory(BaseModel): + model_config = ConfigDict(extra="forbid", validate_by_alias=False, validate_by_name=True) + count_value: int = Field(alias="count") + color_value: str = Field(alias="color") + + calls = 0 + + @define_tool("get_inventory", description="Get the current widget inventory.") + def get_inventory() -> str: + nonlocal calls + calls += 1 + return "The inventory contains 42 red widgets." + + async with await ctx.client.create_session( + **config(ctx), tools=[get_inventory], streaming=True + ) as session: + deltas = [] + unsubscribe = session.on( + lambda event: ( + deltas.append(event.data) + if isinstance(event.data, AssistantMessageDeltaData) + else None + ) + ) + result = await session.send_and_wait_typed( + "Call get_inventory, then report the widget count and color.", + AliasedInventory, + timeout=30, + ) + assert result.count_value == 42 + assert result.color_value == "red" + assert calls > 0 + assert deltas + unsubscribe() + ordinary = await session.send_and_wait( + "Now reply with exactly the plain text HELLO, not JSON.", timeout=30 + ) + assert ordinary.data.content.strip() == "HELLO" + exchanges = await ctx.get_exchanges() + assert len(exchanges) >= 3 + for exchange in exchanges[:-1]: + assert exchange["request"]["response_format"] == { + "type": "json_schema", + "json_schema": { + "name": "response", + "strict": True, + "schema": AliasedInventory.model_json_schema(), + }, + } + assert "response_format" not in exchanges[-1]["request"] + + +async def test_typed_wait_returns_stop_hook_correction(ctx): + stops = 0 + + def stop_hook(_input, _invocation): + nonlocal stops + stops += 1 + if stops == 1: + return { + "decision": "block", + "reason": "Correct the answer to 99, not 42. Do not use tools.", + } + return None + + async with await ctx.client.create_session( + **config(ctx), hooks={"on_agent_stop": stop_hook} + ) as session: + replies = [] + unsubscribe = session.on( + lambda event: ( + replies.append(event.data) + if not event.agent_id and isinstance(event.data, AssistantMessageData) + else None + ) + ) + try: + result = await session.send_and_wait_typed( + "What is 19 + 23? Do not use tools.", Answer, timeout=30 + ) + assert result.answer == 99 + assert stops == 2 + assert [Answer.model_validate_json(reply.content).answer for reply in replies] == [ + 42, + 99, + ] + assert replies[0].originating_message_id + assert replies[0].originating_message_id == replies[1].originating_message_id + finally: + unsubscribe() + + +async def test_send_selects_correlated_response_after_idle(ctx): + entered = asyncio.Event() + release = asyncio.Event() + + @define_tool("read_inventory", description="Read the current widget count and color.") + def read_inventory() -> str: + return "The inventory contains 42 red widgets." + + async def stop_hook(_input, _invocation): + entered.set() + await release.wait() + return None + + async with await ctx.client.create_session( + **config(ctx), tools=[read_inventory], hooks={"on_agent_stop": stop_hook} + ) as session: + idle = asyncio.get_running_loop().create_future() + replies = [] + + def observe(event): + if event.agent_id: + return + if isinstance(event.data, AssistantMessageData): + replies.append(event.data) + elif isinstance(event.data, SessionIdleData) and not idle.done(): + idle.set_result(None) + elif isinstance(event.data, SessionErrorData) and not idle.done(): + idle.set_exception(RuntimeError(event.data.message)) + + unsubscribe = session.on(observe) + try: + origin = await session.send( + "Call read_inventory once, then report the current widget count and color.", + response_schema=Inventory, + ) + await asyncio.wait_for(entered.wait(), 30) + assert not idle.done() + release.set() + await asyncio.wait_for(idle, 30) + final = next( + reply for reply in reversed(replies) if reply.originating_message_id == origin + ) + assert final is replies[-1] + assert not final.tool_requests + assert any(reply.tool_requests for reply in replies) + assert Inventory.model_validate_json(final.content) == Inventory(count=42, color="red") + finally: + release.set() + unsubscribe() + + +async def test_rejects_invalid_formats_before_admission(ctx): + async with await ctx.client.create_session(**config(ctx)) as session: + events = [] + unsubscribe = session.on(events.append) + try: + with pytest.raises(ValueError, match="immediate"): + await session.send_and_wait_typed("Must not be admitted", Answer, mode="immediate") + schema = {"type": "object", "description": "x" * (32 * 1024 * 1024)} + with pytest.raises(Exception, match="32 MiB"): + await session.send_and_wait( + "Must not be admitted", response_schema=schema, timeout=30 + ) + with pytest.raises(Exception, match="32 MiB"): + await session.rpc.send_messages( + SendMessagesRequest( + messages=[], + response_format=ResponseFormat( + type=ResponseFormatType.JSON_SCHEMA, + json_schema=JSONSchemaResponseFormat(name="response", schema=schema), + ), + ) + ) + assert not (await session.rpc.queue.pending_items()).items + assert not any( + isinstance(event.data, (UserMessageData, SessionErrorData)) for event in events + ) + assert not await ctx.get_exchanges() + finally: + unsubscribe() + + +async def test_sends_explicit_schema_for_message_and_batch(ctx): + async with await ctx.client.create_session(**config(ctx)) as session: + idle = asyncio.get_running_loop().create_future() + replies = [] + + def observe(event): + if event.agent_id: + return + if isinstance(event.data, AssistantMessageData): + replies.append(event.data) + elif isinstance(event.data, SessionIdleData) and not idle.done(): + idle.set_result(None) + elif isinstance(event.data, SessionErrorData) and not idle.done(): + idle.set_exception(RuntimeError(event.data.message)) + + unsubscribe = session.on(observe) + schema = Inventory.model_json_schema() + try: + accepted = await session.rpc.send_messages( + SendMessagesRequest( + messages=[ + SendMessageItem(prompt="There are 42 red widgets in stock."), + SendMessageItem(prompt="Report the widget count and color."), + ], + response_format=ResponseFormat( + type=ResponseFormatType.JSON_SCHEMA, + json_schema=JSONSchemaResponseFormat( + name="inventory", schema=schema, strict=True + ), + ), + ) + ) + await asyncio.wait_for(idle, 30) + final = next( + item + for item in reversed(replies) + if item.originating_message_id == accepted.message_ids[-1] + ) + assert Inventory.model_validate_json(final.content) == Inventory(count=42, color="red") + finally: + unsubscribe() + updated = await session.send_and_wait( + "The inventory now has 21 blue widgets. Report the new count and color.", + response_schema=schema, + timeout=30, + ) + assert Inventory.model_validate_json(updated.data.content) == Inventory( + count=21, color="blue" + ) + + +async def test_typed_wait_returns_stop_hook_correction_after_terminal_tool(ctx): + calls = 0 + stops = 0 + + @define_tool( + "lookup_number", + description="Return the number needed for the calculation.", + is_terminal=True, + skip_permission=True, + ) + def lookup_number() -> int: + nonlocal calls + calls += 1 + return 58 + + def stop_hook(_input, _invocation): + nonlocal stops + stops += 1 + if stops == 1: + return { + "decision": "block", + "reason": "Correct the answer to 99, not 63. Do not use tools.", + } + return None + + async with await ctx.client.create_session( + **config(ctx), tools=[lookup_number], hooks={"on_agent_stop": stop_hook} + ) as session: + replies = [] + unsubscribe = session.on( + lambda event: ( + replies.append(event.data) + if not event.agent_id + and isinstance(event.data, AssistantMessageData) + and not event.data.tool_requests + else None + ) + ) + try: + result = await session.send_and_wait_typed( + "Call lookup_number exactly once, then add 5 to the returned number. " + "Do not guess its result.", + Answer, + timeout=30, + ) + assert result.answer == 99 + assert calls == 1 + assert stops == 2 + assert [json.loads(reply.content)["answer"] for reply in replies] == [63, 99] + assert replies[0].originating_message_id == replies[1].originating_message_id + finally: + unsubscribe() + + +async def test_typed_result_after_terminal_tool_and_steering(ctx): + calls = 0 + + @define_tool( + "lookup_number", + description="Return the number needed for the calculation.", + is_terminal=True, + skip_permission=True, + ) + async def lookup_number() -> int: + nonlocal calls + calls += 1 + await session.send( + "Continue with the original calculation. Do not call any more tools.", mode="immediate" + ) + return 58 + + async with await ctx.client.create_session(**config(ctx), tools=[lookup_number]) as session: + result = await session.send_and_wait_typed( + "Call lookup_number exactly once, then add 5 to the returned number. " + "Do not guess its result.", + ToolAnswer, + timeout=30, + ) + assert result == ToolAnswer(answer=63, contract="typed_tool") + assert calls == 1 + exchanges = await ctx.get_exchanges() + assert len(exchanges) >= 2 + for exchange in exchanges[1:]: + assert exchange["request"]["tool_choice"] == "none" + for exchange in exchanges: + assert ( + exchange["request"]["response_format"]["json_schema"]["schema"] + == ToolAnswer.model_json_schema() + ) + + +async def test_typed_wait_returns_late_steering_response(ctx): + stops = 0 + steering_id = None + + async def stop_hook(_input, _invocation): + nonlocal stops, steering_id + stops += 1 + if stops == 1: + steering_id = await session.send( + "Change the answer to 99. Do not use tools.", mode="immediate" + ) + return None + + async with await ctx.client.create_session( + **config(ctx), hooks={"on_agent_stop": stop_hook} + ) as session: + result = await session.send_and_wait_typed( + "What is 19 + 23? Do not use tools.", Answer, timeout=30 + ) + assert result.answer == 99 + assert stops == 2 + assert steering_id + + +async def test_concurrent_typed_sends_return_their_own_results(ctx): + entered = asyncio.Event() + release = asyncio.Event() + + @define_tool("first_number", description="Get the number for the first question.") + async def first_number() -> int: + entered.set() + await release.wait() + return 42 + + async with await ctx.client.create_session(**config(ctx), tools=[first_number]) as session: + first = asyncio.create_task( + session.send_and_wait_typed( + "Call first_number exactly once and report its returned number.", First, timeout=30 + ) + ) + try: + await asyncio.wait_for(entered.wait(), 30) + second = asyncio.create_task( + session.send_and_wait_typed("What is 30 + 7? Do not use tools.", Second, timeout=30) + ) + async with asyncio.timeout(30): + while not (await session.rpc.queue.pending_items()).items: + await asyncio.sleep(0.01) + release.set() + assert (await first).first == 42 + assert (await second).second == 37 + finally: + release.set() + if not first.done(): + first.cancel() + await asyncio.gather(first, return_exceptions=True) diff --git a/python/pyproject.toml b/python/pyproject.toml index e96c587a64..c58b66c67c 100644 --- a/python/pyproject.toml +++ b/python/pyproject.toml @@ -26,7 +26,7 @@ classifiers = [ ] dependencies = [ "python-dateutil>=2.9.0.post0", - "pydantic>=2.0", + "pydantic>=2.11", "httpx>=0.24.0", ] diff --git a/python/test_structured_output.py b/python/test_structured_output.py new file mode 100644 index 0000000000..2b626bc6bf --- /dev/null +++ b/python/test_structured_output.py @@ -0,0 +1,278 @@ +"""Structured-output admission, correlation, and lifecycle tests.""" + +import asyncio +from datetime import UTC, datetime +from unittest.mock import AsyncMock, Mock +from uuid import uuid4 + +import pytest +from pydantic import BaseModel, ConfigDict, Field, ValidationError + +from copilot.session import CopilotSession +from copilot.session_events import SessionEvent + + +class Inventory(BaseModel): + model_config = ConfigDict(extra="forbid") + count: int + color: str + + +def event(kind, data, agent_id=None): + return SessionEvent.from_dict( + { + "id": str(uuid4()), + "timestamp": datetime.now(UTC).isoformat(), + "type": kind, + "data": data, + **({"agentId": agent_id} if agent_id else {}), + } + ) + + +def assistant(content='{"count":42,"color":"red"}', origin="user-1", **extra): + return event( + "assistant.message", + {"messageId": str(uuid4()), "content": content, "originatingMessageId": origin, **extra}, + ) + + +def user(origin="user-1"): + return event("user.message", {"messageId": origin, "content": "inventory"}) + + +def idle(**data): + return event("session.idle", data) + + +def fake_session(events=(), admission_error=None): + client = Mock() + session = CopilotSession("session-1", client) + + async def request(method, params): + assert method == "session.send" + for item in events: + session._dispatch_event(item) + if admission_error: + raise admission_error + return {"messageId": "user-1"} + + client.request = AsyncMock(side_effect=request) + return session, client + + +def assert_clean(session): + assert not session._structured_waits + assert not session._event_handlers + + +@pytest.mark.asyncio +async def test_typed_output_buffers_pre_ack_events_and_infers_schema(): + session, client = fake_session([user(), assistant(), idle()]) + result = await session.send_and_wait_typed("inventory", Inventory, timeout=1) + assert result == Inventory(count=42, color="red") + params = client.request.call_args.args[1] + assert params["responseFormat"] == { + "type": "json_schema", + "jsonSchema": {"name": "response", "strict": True, "schema": Inventory.model_json_schema()}, + } + assert_clean(session) + + +@pytest.mark.asyncio +async def test_typed_output_validates_schema_aliases_including_nested_models(): + class Answer(BaseModel): + model_config = ConfigDict(extra="forbid", validate_by_alias=False, validate_by_name=True) + value: int = Field(alias="answer") + + class Result(BaseModel): + model_config = ConfigDict(extra="forbid", validate_by_alias=False, validate_by_name=True) + value: Answer = Field(alias="result") + values: list[Answer] = Field(alias="results") + + content = '{"result":{"answer":42},"results":[{"answer":99}]}' + session, client = fake_session([assistant(content), idle()]) + result = await session.send_and_wait_typed("answer", Result, timeout=1) + assert result.value.value == 42 + assert result.values[0].value == 99 + schema = client.request.call_args.args[1]["responseFormat"]["jsonSchema"]["schema"] + assert schema == Result.model_json_schema() + assert set(schema["properties"]) == {"result", "results"} + assert set(schema["$defs"]["Answer"]["properties"]) == {"answer"} + assert not Result.model_config["validate_by_alias"] + assert not Answer.model_config["validate_by_alias"] + assert_clean(session) + + +@pytest.mark.asyncio +async def test_raw_schema_is_forwarded_unchanged_and_result_remains_an_event(): + schema = {"type": "object", "properties": {"count": {"type": "integer", "minimum": 2}}} + final = assistant() + session, client = fake_session([final, idle()]) + result = await session.send_and_wait("inventory", response_schema=schema, timeout=1) + assert result is final + assert client.request.call_args.args[1]["responseFormat"]["jsonSchema"]["schema"] == schema + assert "additionalProperties" not in schema + assert_clean(session) + + +@pytest.mark.asyncio +async def test_correlation_tool_commentary_stop_corrections_and_autopilot(): + subagent = assistant('{"count":999,"color":"wrong"}') + subagent.agent_id = "child" + events = [ + idle(), + event("session.error", {"errorType": "test", "message": "before this run"}), + user(), + assistant("working", toolRequests=[{"toolCallId": "tool-1", "name": "inventory"}]), + assistant(), + idle(mode="autopilot"), + assistant('{"count":99,"color":"blue"}'), + subagent, + assistant('{"count":123,"color":"wrong"}', origin="other-user"), + idle(), + ] + session, _ = fake_session(events) + assert await session.send_and_wait_typed("inventory", Inventory, timeout=1) == Inventory( + count=99, color="blue" + ) + assert_clean(session) + + +@pytest.mark.parametrize( + ("events", "error"), + [ + ([user(), idle()], "without a structured"), + ([user(), assistant(" "), idle()], "without a structured"), + ( + [ + assistant(), + assistant("working", toolRequests=[{"toolCallId": "t", "name": "tool"}]), + idle(), + ], + "without a structured", + ), + ([assistant(), idle(aborted=True)], "aborted"), + ( + [user(), event("session.error", {"errorType": "test", "message": "provider failed"})], + "provider failed", + ), + ], +) +@pytest.mark.asyncio +async def test_failed_runs_do_not_return_stale_or_missing_results(events, error): + session, _ = fake_session(events) + with pytest.raises(RuntimeError, match=error): + await session.send_and_wait_typed("inventory", Inventory, timeout=1) + assert_clean(session) + + +@pytest.mark.parametrize( + "content", + [ + "null", + "not JSON", + '{"count":"no","color":"red"}', + '{"count":42}', + '{"count":42,"color":"red","extra":1}', + ], +) +@pytest.mark.asyncio +async def test_typed_output_validates_json(content): + session, _ = fake_session([assistant(content), idle()]) + with pytest.raises(ValidationError): + await session.send_and_wait_typed("inventory", Inventory, timeout=1) + assert_clean(session) + + +@pytest.mark.asyncio +async def test_admission_error_is_not_hidden_by_buffered_success(): + session, _ = fake_session([assistant(), idle()], ValueError("invalid schema")) + with pytest.raises(ValueError, match="invalid schema"): + await session.send_and_wait_typed("inventory", Inventory, timeout=1) + assert_clean(session) + + +@pytest.mark.parametrize("kind", ["timeout", "cancel", "disconnect"]) +@pytest.mark.asyncio +async def test_wait_cleanup(kind): + session, client = fake_session() + admitted = asyncio.Event() + + async def request(method, params): + admitted.set() + return {"messageId": "user-1"} + + client.request.side_effect = request + waiting = asyncio.create_task( + session.send_and_wait_typed( + "inventory", Inventory, timeout=0.05 if kind == "timeout" else 1 + ) + ) + await admitted.wait() + if kind == "cancel": + waiting.cancel() + error = asyncio.CancelledError + elif kind == "disconnect": + session._mark_disconnected() + error = RuntimeError + else: + error = TimeoutError + with pytest.raises(error): + await waiting + assert_clean(session) + + +@pytest.mark.asyncio +async def test_concurrent_typed_waits_keep_their_own_origin(): + session, client = fake_session() + admitted = asyncio.Queue() + + async def request(method, params): + origin = params["prompt"] + admitted.put_nowait(origin) + return {"messageId": origin} + + client.request.side_effect = request + first = asyncio.create_task(session.send_and_wait_typed("first", Inventory, timeout=1)) + second = asyncio.create_task(session.send_and_wait_typed("second", Inventory, timeout=1)) + await admitted.get() + await admitted.get() + session._dispatch_event(assistant(origin="first")) + session._dispatch_event(assistant('{"count":7,"color":"blue"}', origin="second")) + session._dispatch_event(idle()) + assert (await first).count == 42 + assert (await second).count == 7 + assert_clean(session) + + +@pytest.mark.asyncio +async def test_typed_immediate_is_rejected_before_admission(): + session, client = fake_session() + with pytest.raises(ValueError, match="immediate"): + await session.send_and_wait_typed("inventory", Inventory, mode="immediate") + client.request.assert_not_called() + assert_clean(session) + + +@pytest.mark.asyncio +async def test_disconnect_while_admission_is_pending_cancels_rpc_wait(): + session, client = fake_session() + entered = asyncio.Event() + cancelled = asyncio.Event() + + async def request(method, params): + entered.set() + try: + await asyncio.Future() + finally: + cancelled.set() + + client.request.side_effect = request + waiting = asyncio.create_task(session.send_and_wait_typed("inventory", Inventory)) + await entered.wait() + session._mark_disconnected() + with pytest.raises(RuntimeError, match="Session closed"): + await asyncio.wait_for(waiting, 1) + assert cancelled.is_set() + assert_clean(session) diff --git a/rust/README.md b/rust/README.md index 11d9637b22..bde9c26215 100644 --- a/rust/README.md +++ b/rust/README.md @@ -873,7 +873,49 @@ session .await?; ``` -Default timeout is 60 seconds. Only one `send_and_wait` can be active per session — concurrent calls return an error. +Default timeout is 60 seconds. Only one unformatted `send_and_wait` can be active +per session; it also prevents other sends until it completes. + +### Structured output (experimental) + +Enable the existing `derive` feature and use the same `schemars`/Serde integration +as typed custom tools: + +```rust,no_run +# #[cfg(feature = "derive")] +# mod example { +use schemars::JsonSchema; +use serde::Deserialize; + +#[derive(Deserialize, JsonSchema)] +#[serde(deny_unknown_fields)] +struct Inventory { + count: i32, + color: String, +} + +# async fn example(session: &github_copilot_sdk::session::Session) -> Result<(), github_copilot_sdk::Error> { +let inventory: Inventory = session + .send_and_wait_typed("Call get_inventory, then report the widget count and color.") + .await?; +# Ok(()) +# } +# } +``` + +The helper uses the existing `schema_for::()` generator and deserializes the +final JSON. Serde deserialization is not full JSON Schema validation. Provider +schema restrictions apply; `deny_unknown_fields` closes objects for strict output. +For explicit schemas, `MessageOptions::with_response_schema` works with `send` or +`send_and_wait` without the `derive` feature and returns ordinary events. + +Schemas apply to one run, including tools, steering, and stop-hook corrections, +not independent sends or subagents. Streaming remains text. Structured waits +select the last correlated root message without tool requests at non-autopilot +idle and support concurrent structured waits with independent results. Later +queued work can delay idle. Aborts, session errors after the run starts, missing +output, and event-stream lag fail the wait. Dropping the future or timing out +unsubscribes without aborting the agent. Immediate steering cannot set a schema. ### Newtypes diff --git a/rust/src/session.rs b/rust/src/session.rs index 82e97b6606..4ecd7e457a 100644 --- a/rust/src/session.rs +++ b/rust/src/session.rs @@ -144,6 +144,90 @@ struct IdleWaiter { first_assistant_message_seen: bool, } +fn structured_output_error(message: impl Into) -> Error { + Error::with_message( + ErrorKind::Session(SessionErrorKind::AgentError), + message.into(), + ) +} + +fn is_structured_output_event(event: &SessionEvent) -> bool { + matches!( + event.parsed_type(), + SessionEventType::UserMessage + | SessionEventType::AssistantMessage + | SessionEventType::SessionIdle + | SessionEventType::SessionError + ) +} + +struct StructuredOutputState { + message_id: String, + started: bool, + final_message: Option, +} + +impl StructuredOutputState { + fn observe(&mut self, event: SessionEvent) -> Result, Error> { + if event.agent_id.as_deref().is_some_and(|id| !id.is_empty()) { + return Ok(None); + } + match event.parsed_type() { + SessionEventType::UserMessage => { + let data: crate::session_events::UserMessageData = + serde_json::from_value(event.data)?; + if data.message_id.as_deref() == Some(self.message_id.as_str()) { + self.started = true; + } + } + SessionEventType::AssistantMessage => { + let data: crate::session_events::AssistantMessageData = + serde_json::from_value(event.data.clone())?; + if data.originating_message_id.as_deref() == Some(self.message_id.as_str()) { + self.started = true; + self.final_message = + if data.tool_requests.is_some_and(|tools| !tools.is_empty()) { + None + } else { + Some(event) + }; + } + } + SessionEventType::SessionIdle if self.started => { + let data: SessionIdleData = serde_json::from_value(event.data)?; + if data.mode == Some(SessionMode::Autopilot) { + return Ok(None); + } + if data.aborted == Some(true) { + return Err(structured_output_error( + "session aborted before structured output completed", + )); + } + let result = self.final_message.take().ok_or_else(|| { + structured_output_error("run completed without a structured assistant response") + })?; + let data: crate::session_events::AssistantMessageData = + serde_json::from_value(result.data.clone())?; + if data.content.trim().is_empty() { + return Err(structured_output_error( + "run completed without a structured assistant response", + )); + } + return Ok(Some(result)); + } + SessionEventType::SessionError if self.started => { + let data: SessionErrorData = serde_json::from_value(event.data)?; + return Err(structured_output_error(format!( + "session error: {}", + data.message + ))); + } + _ => {} + } + Ok(None) + } +} + /// RAII guard that clears the [`Session::idle_waiter`] slot on drop. Used /// by [`Session::send_and_wait`] to ensure the slot doesn't leak if the /// caller's future is cancelled (outer `tokio::time::timeout` / `select!` @@ -543,6 +627,12 @@ impl Session { if let Some(display_prompt) = opts.display_prompt { params["displayPrompt"] = serde_json::to_value(display_prompt)?; } + if let Some(schema) = opts.response_schema { + params["responseFormat"] = serde_json::json!({ + "type": "json_schema", + "jsonSchema": { "name": "response", "strict": true, "schema": schema } + }); + } let trace_ctx = if opts.traceparent.is_some() || opts.tracestate.is_some() { TraceContext { traceparent: opts.traceparent, @@ -576,9 +666,11 @@ impl Session { /// returning the last `assistant.message` event captured during streaming. /// Times out after `MessageOptions::wait_timeout` (default 60 seconds). /// - /// Only one `send_and_wait` call may be active per session at a time. - /// Calling [`send`](Self::send) while a `send_and_wait` - /// is in flight will also return an error. + /// Only one unformatted `send_and_wait` may be active per session. Calling + /// [`send`](Self::send) during that wait also returns an error. Schema-bearing + /// waits instead correlate by originating message ID and support concurrency. + /// They select the last root message without tool requests at non-autopilot + /// idle, failing on aborted idle, session errors after starting, or no result. /// /// # Cancel safety /// @@ -593,6 +685,9 @@ impl Session { ) -> Result, Error> { let total_start = Instant::now(); let opts = opts.into(); + if opts.response_schema.is_some() { + return self.send_and_wait_structured(opts).await.map(Some); + } let timeout_duration = opts.wait_timeout.unwrap_or(Duration::from_secs(60)); let (tx, rx) = oneshot::channel(); @@ -648,6 +743,82 @@ impl Session { } } + /// Infer an output schema with the same `schemars` integration as custom tools, + /// then deserialize the final correlated root response at non-autopilot idle. + /// + /// Requires the `derive` feature. Provider schema restrictions apply. Serde + /// validates JSON/type compatibility, not every JSON Schema constraint. + /// Options must not specify a schema or immediate delivery. Dropping this + /// future or timing out unsubscribes the wait without aborting agent work. + #[cfg(feature = "derive")] + pub async fn send_and_wait_typed(&self, opts: impl Into) -> Result + where + T: schemars::JsonSchema + serde::de::DeserializeOwned, + { + let mut opts = opts.into(); + if opts.response_schema.is_some() + || opts.mode == Some(crate::types::DeliveryMode::Immediate) + { + return Err(Error::with_message( + ErrorKind::InvalidConfig, + "typed structured output cannot specify a response schema or immediate delivery", + )); + } + opts.response_schema = Some(crate::tool::schema_for::()); + let event = self.send_and_wait_structured(opts).await?; + let data: crate::session_events::AssistantMessageData = serde_json::from_value(event.data)?; + serde_json::from_str::>(&data.content)?.ok_or_else(|| { + structured_output_error("structured response was JSON null, not a result") + }) + } + + async fn send_and_wait_structured(&self, opts: MessageOptions) -> Result { + let duration = opts.wait_timeout.unwrap_or(Duration::from_secs(60)); + let mut events = self.subscribe(); + let wait = async { + let mut admission = Box::pin(self.send(opts)); + let mut pending = Vec::new(); + let message_id = loop { + tokio::select! { + result = &mut admission => break result?, + event = events.recv() => { + let event = event.map_err(|err| structured_output_error(err.to_string()))?; + if is_structured_output_event(&event) { + pending.push(event); + } + } + _ = self.shutdown.cancelled() => + return Err(structured_output_error("session closed before structured output completed")), + } + }; + let mut state = StructuredOutputState { + message_id, + started: false, + final_message: None, + }; + for event in pending { + if let Some(result) = state.observe(event)? { + return Ok(result); + } + } + loop { + tokio::select! { + event = events.recv() => { + let event = event.map_err(|err| structured_output_error(err.to_string()))?; + if let Some(result) = state.observe(event)? { + return Ok(result); + } + } + _ = self.shutdown.cancelled() => + return Err(structured_output_error("session closed before structured output completed")), + } + } + }; + tokio::time::timeout(duration, wait) + .await + .map_err(|_| Error::from(ErrorKind::Session(SessionErrorKind::Timeout(duration))))? + } + /// Retrieve the session's timeline events. pub async fn get_events(&self) -> Result, Error> { let result = self diff --git a/rust/src/types.rs b/rust/src/types.rs index 6e3af9273f..131dd24570 100644 --- a/rust/src/types.rs +++ b/rust/src/types.rs @@ -5505,6 +5505,9 @@ pub enum AgentMode { #[derive(Debug, Clone)] #[non_exhaustive] pub struct MessageOptions { + /// Per-run JSON Schema. Independent sends and subagents do not inherit it. + /// Immediate steering must not specify a schema. Streaming events remain text. + pub response_schema: Option, /// The user prompt to send. pub prompt: String, /// Optional message provenance. When `None`, the field is omitted, @@ -5549,6 +5552,7 @@ impl MessageOptions { pub fn new(prompt: impl Into) -> Self { Self { prompt: prompt.into(), + response_schema: None, source: None, mode: None, agent_mode: None, @@ -5567,6 +5571,12 @@ impl MessageOptions { self } + /// Request provider-native structured output for this run. + pub fn with_response_schema(mut self, schema: Value) -> Self { + self.response_schema = Some(schema); + self + } + /// Set the message delivery mode for this turn. /// /// Pass [`DeliveryMode::Immediate`] to interrupt the session and run diff --git a/rust/tests/e2e.rs b/rust/tests/e2e.rs index 9d1c868fe9..cae2f92eb9 100644 --- a/rust/tests/e2e.rs +++ b/rust/tests/e2e.rs @@ -132,6 +132,9 @@ mod session_todos_changed; mod skills; #[path = "e2e/streaming_fidelity.rs"] mod streaming_fidelity; +#[cfg(feature = "derive")] +#[path = "e2e/structured_output.rs"] +mod structured_output; #[path = "e2e/subagent_hooks.rs"] mod subagent_hooks; #[path = "e2e/support.rs"] diff --git a/rust/tests/e2e/structured_output.rs b/rust/tests/e2e/structured_output.rs new file mode 100644 index 0000000000..7fdbce4fc3 --- /dev/null +++ b/rust/tests/e2e/structured_output.rs @@ -0,0 +1,586 @@ +use std::collections::HashMap; +use std::sync::Arc; +use std::sync::atomic::{AtomicUsize, Ordering}; +use std::time::Duration; + +use async_trait::async_trait; +use github_copilot_sdk::handler::ApproveAllHandler; +use github_copilot_sdk::hooks::{AgentStopInput, AgentStopOutput, HookContext, SessionHooks}; +use github_copilot_sdk::rpc::SendMessagesRequest; +use github_copilot_sdk::tool::{define_tool, schema_for}; +use github_copilot_sdk::{DeliveryMode, MessageOptions, ProviderConfig, SessionConfig, ToolResult}; +use schemars::JsonSchema; +use serde::Deserialize; +use serde_json::json; +use tokio::sync::Semaphore; + +use super::support::{DEFAULT_TEST_TOKEN, assistant_message_content, collect_until_idle}; + +#[derive(Debug, PartialEq, Deserialize, JsonSchema)] +#[serde(deny_unknown_fields)] +struct Inventory { + count: i32, + color: String, +} + +#[derive(Deserialize, JsonSchema)] +#[serde(deny_unknown_fields)] +struct Answer { + answer: i32, +} + +#[derive(Debug, PartialEq, Deserialize, JsonSchema)] +#[serde(deny_unknown_fields)] +struct ToolAnswer { + answer: i32, + contract: String, +} + +#[derive(Deserialize, JsonSchema)] +#[serde(deny_unknown_fields)] +struct First { + first: i32, +} + +#[derive(Deserialize, JsonSchema)] +#[serde(deny_unknown_fields)] +struct Second { + second: i32, +} + +#[derive(Deserialize, JsonSchema)] +struct NoArgs {} + +fn config(proxy: &str) -> SessionConfig { + SessionConfig::default() + .with_model("gpt-4.1") + .with_permission_handler(Arc::new(ApproveAllHandler)) + .with_available_tools(Vec::::new()) + .with_provider( + ProviderConfig::new(proxy) + .with_provider_type("openai") + .with_wire_api("completions") + .with_api_key(DEFAULT_TEST_TOKEN) + .with_model_id("gpt-4.1") + .with_wire_model("gpt-4.1") + .with_headers(HashMap::from([ + ( + "Copilot-Integration-Id".into(), + "copilot-developer-cli".into(), + ), + ("Copilot-Harness-Id".into(), "copilot-sdk".into()), + ("X-GitHub-Api-Version".into(), "2026-08-01".into()), + ])), + ) +} + +#[tokio::test] +async fn infers_typed_result_after_custom_tool() { + super::support::with_shared_e2e_context( + &E2E, + "structured_output", + "infers_typed_result_after_custom_tool", + |ctx| { + Box::pin(async move { + let calls = Arc::new(AtomicUsize::new(0)); + let counter = calls.clone(); + let tool = define_tool( + "get_inventory", + "Get the current widget inventory.", + move |_inv, _: NoArgs| { + counter.fetch_add(1, Ordering::SeqCst); + async { + Ok(ToolResult::Text( + "The inventory contains 42 red widgets.".into(), + )) + } + }, + ); + let client = ctx.start_client().await; + let session = client + .create_session( + config(ctx.proxy_url()) + .with_tools(vec![tool]) + .with_streaming(true), + ) + .await + .unwrap(); + let events = session.subscribe(); + let result: Inventory = session + .send_and_wait_typed( + "Call get_inventory, then report the widget count and color.", + ) + .await + .unwrap(); + assert_eq!( + result, + Inventory { + count: 42, + color: "red".into() + } + ); + assert!(calls.load(Ordering::SeqCst) > 0); + assert!( + collect_until_idle(events).await.iter().any(|event| event.event_type == "assistant.message_delta"), + "Typed wait must preserve streaming text updates" + ); + let ordinary = session + .send_and_wait("Now reply with exactly the plain text HELLO, not JSON.") + .await + .unwrap() + .unwrap(); + assert_eq!(assistant_message_content(&ordinary).trim(), "HELLO"); + let exchanges = ctx.exchanges(); + assert!(exchanges.len() >= 3); + for exchange in &exchanges[..exchanges.len() - 1] { + assert_eq!(exchange["request"]["response_format"], json!({ + "type": "json_schema", + "json_schema": {"name": "response", "strict": true, "schema": schema_for::()} + })); + } + assert!(exchanges.last().unwrap()["request"].get("response_format").is_none()); + session.disconnect().await.unwrap(); + client.stop().await.unwrap(); + }) + }, + ) + .await; +} + +#[tokio::test] +async fn sends_explicit_schema_for_message_and_batch() { + super::support::with_shared_e2e_context(&E2E, "structured_output", "sends_explicit_schema_for_message_and_batch", |ctx| { + Box::pin(async move { + let client = ctx.start_client().await; + let session = client.create_session(config(ctx.proxy_url())).await.unwrap(); + let mut events = session.subscribe(); + let request: SendMessagesRequest = serde_json::from_value(json!({ + "messages": [ + {"prompt":"There are 42 red widgets in stock."}, + {"prompt":"Report the widget count and color."} + ], + "responseFormat": { + "type":"json_schema", + "jsonSchema":{"name":"inventory","strict":true,"schema":schema_for::()} + } + })).unwrap(); + let accepted = session.rpc().send_messages(request).await.unwrap(); + let mut final_message = None; + loop { + let event = tokio::time::timeout(Duration::from_secs(30), events.recv()).await.unwrap().unwrap(); + if event.agent_id.is_some() { continue; } + match event.event_type.as_str() { + "assistant.message" if event.data["originatingMessageId"] == *accepted.message_ids.last().unwrap() => final_message = Some(event), + "session.idle" => break, + "session.error" => panic!("session error: {:?}", event.data), + _ => {} + } + } + let batch: Inventory = serde_json::from_str(&assistant_message_content(&final_message.unwrap())).unwrap(); + assert_eq!(batch, Inventory { count: 42, color: "red".into() }); + let raw = session.send_and_wait(MessageOptions::new( + "The inventory now has 21 blue widgets. Report the new count and color." + ).with_response_schema(schema_for::())).await.unwrap().unwrap(); + let updated: Inventory = serde_json::from_str(&assistant_message_content(&raw)).unwrap(); + assert_eq!(updated, Inventory { count: 21, color: "blue".into() }); + session.disconnect().await.unwrap(); + client.stop().await.unwrap(); + }) + }).await; +} + +struct CorrectionHook { + calls: AtomicUsize, + reason: &'static str, +} + +#[async_trait] +impl SessionHooks for CorrectionHook { + async fn on_agent_stop( + &self, + _input: AgentStopInput, + _ctx: HookContext, + ) -> Option { + (self.calls.fetch_add(1, Ordering::SeqCst) == 0).then(|| AgentStopOutput { + decision: Some("block".into()), + reason: Some(self.reason.into()), + }) + } +} + +#[tokio::test] +async fn typed_wait_returns_stop_hook_correction_after_terminal_tool() { + super::support::with_shared_e2e_context(&E2E, "structured_output", "typed_wait_returns_stop_hook_correction_after_terminal_tool", |ctx| { + Box::pin(async move { + let calls = Arc::new(AtomicUsize::new(0)); + let counter = calls.clone(); + let mut tool = define_tool("lookup_number", "Return the number needed for the calculation.", move |_inv, _: NoArgs| { + counter.fetch_add(1, Ordering::SeqCst); + async { Ok(ToolResult::Text("58".into())) } + }); + tool.is_terminal = true; + tool.skip_permission = true; + let hooks = Arc::new(CorrectionHook { + calls: AtomicUsize::new(0), + reason: "Correct the answer to 99, not 63. Do not use tools.", + }); + let client = ctx.start_client().await; + let session = client.create_session(config(ctx.proxy_url()).with_tools(vec![tool]).with_hooks(hooks.clone())).await.unwrap(); + let result: Answer = session.send_and_wait_typed( + "Call lookup_number exactly once, then add 5 to the returned number. Do not guess its result." + ).await.unwrap(); + assert_eq!(result.answer, 99); + assert_eq!(calls.load(Ordering::SeqCst), 1); + assert_eq!(hooks.calls.load(Ordering::SeqCst), 2); + session.disconnect().await.unwrap(); + client.stop().await.unwrap(); + }) + }).await; +} + +#[tokio::test] +async fn concurrent_typed_sends_return_their_own_results() { + super::support::with_shared_e2e_context( + &E2E, + "structured_output", + "concurrent_typed_sends_return_their_own_results", + |ctx| { + Box::pin(async move { + let entered = Arc::new(Semaphore::new(0)); + let release = Arc::new(Semaphore::new(0)); + let tool_entered = entered.clone(); + let tool_release = release.clone(); + let tool = define_tool( + "first_number", + "Get the number for the first question.", + move |_inv, _: NoArgs| { + let entered = tool_entered.clone(); + let release = tool_release.clone(); + async move { + entered.add_permits(1); + release.acquire().await.unwrap().forget(); + Ok(ToolResult::Text("42".into())) + } + }, + ); + let client = ctx.start_client().await; + let session = client + .create_session(config(ctx.proxy_url()).with_tools(vec![tool])) + .await + .unwrap(); + let first = session.send_and_wait_typed::( + "Call first_number exactly once and report its returned number.", + ); + let second = async { + entered.acquire().await.unwrap().forget(); + let waiting = + session.send_and_wait_typed::("What is 30 + 7? Do not use tools."); + let release_queued = async { + loop { + let queue = session.rpc().queue().pending_items().await.unwrap(); + if !queue.items.is_empty() { + break; + } + tokio::time::sleep(Duration::from_millis(10)).await; + } + release.add_permits(1); + }; + let (result, ()) = tokio::join!(waiting, release_queued); + result + }; + let (first, second) = tokio::join!(first, second); + assert_eq!(first.unwrap().first, 42); + assert_eq!(second.unwrap().second, 37); + session.disconnect().await.unwrap(); + client.stop().await.unwrap(); + }) + }, + ) + .await; +} + +#[tokio::test] +async fn typed_wait_returns_stop_hook_correction() { + super::support::with_shared_e2e_context( + &E2E, + "structured_output", + "typed_wait_returns_stop_hook_correction", + |ctx| { + Box::pin(async move { + let hooks = Arc::new(CorrectionHook { + calls: AtomicUsize::new(0), + reason: "Correct the answer to 99, not 42. Do not use tools.", + }); + let client = ctx.start_client().await; + let session = client + .create_session(config(ctx.proxy_url()).with_hooks(hooks.clone())) + .await + .unwrap(); + let events = session.subscribe(); + let result: Answer = session + .send_and_wait_typed("What is 19 + 23? Do not use tools.") + .await + .unwrap(); + assert_eq!(result.answer, 99); + assert_eq!(hooks.calls.load(Ordering::SeqCst), 2); + let replies: Vec<_> = collect_until_idle(events) + .await + .into_iter() + .filter(|event| { + event.agent_id.is_none() && event.event_type == "assistant.message" + }) + .collect(); + assert_eq!(replies.len(), 2); + assert!( + !replies[0].data["originatingMessageId"] + .as_str() + .unwrap() + .is_empty() + ); + assert_eq!( + replies[0].data["originatingMessageId"], + replies[1].data["originatingMessageId"] + ); + for (reply, expected) in replies.iter().zip([42, 99]) { + assert_eq!( + serde_json::from_str::(&assistant_message_content(reply)) + .unwrap() + .answer, + expected + ); + } + session.disconnect().await.unwrap(); + client.stop().await.unwrap(); + }) + }, + ) + .await; +} + +struct StopGate { + entered: Semaphore, + release: Semaphore, + calls: AtomicUsize, +} + +impl StopGate { + fn new() -> Self { + Self { + entered: Semaphore::new(0), + release: Semaphore::new(0), + calls: AtomicUsize::new(0), + } + } +} + +#[async_trait] +impl SessionHooks for StopGate { + async fn on_agent_stop( + &self, + _input: AgentStopInput, + _ctx: HookContext, + ) -> Option { + if self.calls.fetch_add(1, Ordering::SeqCst) == 0 { + self.entered.add_permits(1); + tokio::time::timeout(Duration::from_secs(30), self.release.acquire()) + .await + .unwrap() + .unwrap() + .forget(); + } + None + } +} + +#[tokio::test] +async fn typed_wait_returns_late_steering_response() { + super::support::with_shared_e2e_context( + &E2E, + "structured_output", + "typed_wait_returns_late_steering_response", + |ctx| { + Box::pin(async move { + let gate = Arc::new(StopGate::new()); + let client = ctx.start_client().await; + let session = client + .create_session(config(ctx.proxy_url()).with_hooks(gate.clone())) + .await + .unwrap(); + let events = session.subscribe(); + let waiting = + session.send_and_wait_typed::("What is 19 + 23? Do not use tools."); + let steer = async { + tokio::time::timeout(Duration::from_secs(30), gate.entered.acquire()) + .await + .unwrap() + .unwrap() + .forget(); + let origin = session + .send( + MessageOptions::new("Change the answer to 99. Do not use tools.") + .with_mode(DeliveryMode::Immediate), + ) + .await + .unwrap(); + gate.release.add_permits(1); + origin + }; + let (result, steering_id) = tokio::join!(waiting, steer); + assert_eq!(result.unwrap().answer, 99); + assert_eq!(gate.calls.load(Ordering::SeqCst), 2); + let replies: Vec<_> = collect_until_idle(events) + .await + .into_iter() + .filter(|event| { + event.agent_id.is_none() && event.event_type == "assistant.message" + }) + .collect(); + assert_eq!(replies.len(), 2); + assert_ne!(replies[0].data["originatingMessageId"], steering_id); + assert!( + !replies[0].data["originatingMessageId"] + .as_str() + .unwrap() + .is_empty() + ); + assert_eq!( + replies[0].data["originatingMessageId"], + replies[1].data["originatingMessageId"] + ); + for (reply, expected) in replies.iter().zip([42, 99]) { + assert_eq!( + serde_json::from_str::(&assistant_message_content(reply)) + .unwrap() + .answer, + expected + ); + } + session.disconnect().await.unwrap(); + client.stop().await.unwrap(); + }) + }, + ) + .await; +} + +#[tokio::test] +async fn send_selects_correlated_response_after_idle() { + super::support::with_shared_e2e_context( + &E2E, "structured_output", "send_selects_correlated_response_after_idle", |ctx| { + Box::pin(async move { + let gate = Arc::new(StopGate::new()); + let tool = define_tool("read_inventory", "Read the current widget count and color.", + |_inv, _: NoArgs| async { Ok(ToolResult::Text("The inventory contains 42 red widgets.".into())) }); + let client = ctx.start_client().await; + let session = client.create_session(config(ctx.proxy_url()).with_hooks(gate.clone()).with_tools(vec![tool])).await.unwrap(); + let events = session.subscribe(); + let origin = session.send(MessageOptions::new("Call read_inventory once, then report the current widget count and color.") + .with_response_schema(schema_for::())).await.unwrap(); + let collecting = collect_until_idle(events); + tokio::pin!(collecting); + tokio::select! { + biased; + _ = &mut collecting => panic!("Idle must wait for the stop hook"), + permit = tokio::time::timeout(Duration::from_secs(30), gate.entered.acquire()) => permit.unwrap().unwrap().forget(), + } + gate.release.add_permits(1); + let replies: Vec<_> = collecting.await.into_iter() + .filter(|event| event.agent_id.is_none() && event.event_type == "assistant.message") + .collect(); + let reply = replies.last().unwrap(); + assert_eq!(reply.data["originatingMessageId"], origin); + assert!(reply.data["toolRequests"].as_array().is_none_or(|tools| tools.is_empty())); + assert!(replies.iter().any(|reply| reply.data["toolRequests"].as_array().is_some_and(|tools| !tools.is_empty()))); + assert_eq!(serde_json::from_str::(&assistant_message_content(reply)).unwrap(), Inventory { count: 42, color: "red".into() }); + session.disconnect().await.unwrap(); + client.stop().await.unwrap(); + }) + }, + ).await; +} + +#[tokio::test] +async fn rejects_invalid_formats_before_admission() { + super::support::with_shared_e2e_context( + &E2E, "structured_output", "rejects_invalid_formats_before_admission", |ctx| { + Box::pin(async move { + let client = ctx.start_client().await; + let session = client.create_session(config(ctx.proxy_url())).await.unwrap(); + let immediate = session.send_and_wait_typed::( + MessageOptions::new("Must not be admitted").with_mode(DeliveryMode::Immediate) + ).await.err().unwrap(); + assert!(immediate.to_string().contains("immediate")); + let schema = json!({"type": "object", "description": "x".repeat(32 * 1024 * 1024)}); + let oversized = session.send_and_wait(MessageOptions::new("Must not be admitted") + .with_response_schema(schema.clone())).await.err().unwrap(); + assert!(oversized.to_string().contains("32 MiB"), "{oversized}"); + let request: SendMessagesRequest = serde_json::from_value(json!({ + "messages": [], + "responseFormat": {"type": "json_schema", "jsonSchema": {"name": "response", "schema": schema}} + })).unwrap(); + let batch = session.rpc().send_messages(request).await.err().unwrap(); + assert!(batch.to_string().contains("32 MiB"), "{batch}"); + assert!(session.rpc().queue().pending_items().await.unwrap().items.is_empty()); + for event in session.get_events().await.unwrap() { + assert_ne!(event.event_type, "user.message"); + assert_ne!(event.event_type, "session.error"); + } + assert!(ctx.exchanges().is_empty()); + session.disconnect().await.unwrap(); + client.stop().await.unwrap(); + }) + }, + ).await; +} + +#[tokio::test] +async fn typed_result_after_terminal_tool_and_steering() { + super::support::with_shared_e2e_context( + &E2E, "structured_output", "typed_result_after_terminal_tool_and_steering", |ctx| { + Box::pin(async move { + let entered = Arc::new(Semaphore::new(0)); + let release = Arc::new(Semaphore::new(0)); + let calls = Arc::new(AtomicUsize::new(0)); + let (tool_entered, tool_release, tool_calls) = (entered.clone(), release.clone(), calls.clone()); + let mut tool = define_tool("lookup_number", "Return the number needed for the calculation.", + move |_inv, _: NoArgs| { + let (entered, release, calls) = (tool_entered.clone(), tool_release.clone(), tool_calls.clone()); + async move { + calls.fetch_add(1, Ordering::SeqCst); + entered.add_permits(1); + tokio::time::timeout(Duration::from_secs(30), release.acquire()).await.unwrap().unwrap().forget(); + Ok(ToolResult::Text("58".into())) + } + }); + tool.is_terminal = true; + tool.skip_permission = true; + let client = ctx.start_client().await; + let session = client.create_session(config(ctx.proxy_url()).with_tools(vec![tool])).await.unwrap(); + let waiting = session.send_and_wait_typed::( + "Call lookup_number exactly once, then add 5 to the returned number. Do not guess its result."); + let steer = async { + tokio::time::timeout(Duration::from_secs(30), entered.acquire()).await.unwrap().unwrap().forget(); + session.send(MessageOptions::new("Continue with the original calculation. Do not call any more tools.") + .with_mode(DeliveryMode::Immediate)).await.unwrap(); + release.add_permits(1); + }; + let (result, ()) = tokio::join!(waiting, steer); + assert_eq!(result.unwrap(), ToolAnswer { answer: 63, contract: "typed_tool".into() }); + assert_eq!(calls.load(Ordering::SeqCst), 1); + let exchanges = ctx.exchanges(); + assert!(exchanges.len() >= 2); + for exchange in &exchanges[1..] { + assert_eq!(exchange["request"]["tool_choice"], "none"); + } + for exchange in &exchanges { + assert_eq!(exchange["request"]["response_format"]["json_schema"]["schema"], schema_for::()); + } + session.disconnect().await.unwrap(); + client.stop().await.unwrap(); + }) + }, + ).await; +} + +static E2E: super::support::SharedE2eGroup = + super::support::SharedE2eGroup::standard("structured_output", 9); diff --git a/rust/tests/session_test.rs b/rust/tests/session_test.rs index 0a8bf52d5a..c04e9e53c8 100644 --- a/rust/tests/session_test.rs +++ b/rust/tests/session_test.rs @@ -3973,6 +3973,306 @@ async fn send_and_wait_returns_last_assistant_message_on_idle() { assert_eq!(event.data["message"], "Hello back!"); } +#[cfg(feature = "derive")] +#[derive(Debug, PartialEq, serde::Deserialize, schemars::JsonSchema)] +#[serde(deny_unknown_fields)] +struct StructuredInventory { + count: i32, + color: String, +} + +#[cfg(feature = "derive")] +#[tokio::test] +async fn structured_output_infers_schema_and_buffers_pre_ack_corrections() { + let (session, mut server) = create_session_pair().await; + let waiting = tokio::spawn(async move { + session + .send_and_wait_typed::("inventory") + .await + }); + let request = server.read_request().await; + let schema = &request["params"]["responseFormat"]["jsonSchema"]["schema"]; + assert_eq!(schema["properties"]["count"]["type"], "integer"); + assert_eq!(schema["additionalProperties"], false); + server + .send_event("session.idle", serde_json::json!({})) + .await; + for (origin, count) in [("user-1", 42), ("user-1", 99), ("other", 123)] { + server + .send_event( + "assistant.message", + serde_json::json!({ + "messageId": "assistant", "originatingMessageId": origin, + "content": format!(r#"{{"count":{count},"color":"red"}}"#) + }), + ) + .await; + if count == 42 { + server + .send_event("session.idle", serde_json::json!({"mode":"autopilot"})) + .await; + } + } + server + .send_event("session.idle", serde_json::json!({})) + .await; + server + .respond(&request, serde_json::json!({"messageId":"user-1"})) + .await; + assert_eq!( + timeout(TIMEOUT, waiting).await.unwrap().unwrap().unwrap(), + StructuredInventory { + count: 99, + color: "red".into() + } + ); +} + +#[cfg(feature = "derive")] +#[tokio::test] +async fn structured_output_preserves_wide_integers() { + #[derive(Debug, PartialEq, serde::Deserialize, schemars::JsonSchema)] + struct WideIntegers { + signed: i128, + unsigned: u128, + } + + for signed in [i128::MIN, i128::MAX, 18446744073709551616] { + let (session, mut server) = create_session_pair().await; + let waiting = tokio::spawn(async move { + session + .send_and_wait_typed::("wide integers") + .await + }); + let request = server.read_request().await; + let content = format!(r#"{{"signed":{signed},"unsigned":{}}}"#, u128::MAX); + server + .respond(&request, serde_json::json!({"messageId":"user-1"})) + .await; + server + .send_event( + "assistant.message", + serde_json::json!({ + "messageId":"assistant", "originatingMessageId":"user-1", "content":content + }), + ) + .await; + server + .send_event("session.idle", serde_json::json!({})) + .await; + assert_eq!( + timeout(TIMEOUT, waiting).await.unwrap().unwrap().unwrap(), + WideIntegers { + signed, + unsigned: u128::MAX + } + ); + } +} + +#[cfg(feature = "derive")] +#[tokio::test] +async fn structured_output_rejects_null_even_for_optional_results() { + let (session, mut server) = create_session_pair().await; + let waiting = tokio::spawn(async move { + session + .send_and_wait_typed::>("inventory") + .await + }); + let request = server.read_request().await; + server + .respond(&request, serde_json::json!({"messageId":"user-1"})) + .await; + server + .send_event( + "assistant.message", + serde_json::json!({ + "messageId":"assistant", "originatingMessageId":"user-1", "content":" \nnull\t " + }), + ) + .await; + server + .send_event("session.idle", serde_json::json!({})) + .await; + let error = timeout(TIMEOUT, waiting) + .await + .unwrap() + .unwrap() + .unwrap_err(); + assert!(error.to_string().contains("JSON null"), "{error}"); +} + +#[cfg(feature = "derive")] +#[tokio::test] +async fn structured_output_rejects_invalid_results() { + for content in [ + "null", + "not JSON", + r#"{"count":"bad","color":"red"}"#, + r#"{"count":42}"#, + ] { + let (session, mut server) = create_session_pair().await; + let waiting = tokio::spawn(async move { + session + .send_and_wait_typed::("inventory") + .await + }); + let request = server.read_request().await; + server + .respond(&request, serde_json::json!({"messageId":"user-1"})) + .await; + server + .send_event( + "assistant.message", + serde_json::json!({ + "messageId":"assistant", "originatingMessageId":"user-1", "content":content + }), + ) + .await; + server + .send_event("session.idle", serde_json::json!({})) + .await; + assert!( + timeout(TIMEOUT, waiting).await.unwrap().unwrap().is_err(), + "{content}" + ); + } +} + +#[tokio::test] +async fn structured_output_raw_schema_and_concurrent_waits_are_independent() { + let (session, mut server) = create_session_pair().await; + let session = Arc::new(session); + let mut waits = Vec::new(); + for origin in ["first", "second"] { + let session = session.clone(); + waits.push(tokio::spawn(async move { + session + .send_and_wait(MessageOptions::new(origin).with_response_schema( + serde_json::json!({"type":"object","description":"unchanged"}), + )) + .await + })); + let request = server.read_request().await; + assert_eq!( + request["params"]["responseFormat"]["jsonSchema"]["schema"], + serde_json::json!({"type":"object","description":"unchanged"}) + ); + server + .respond(&request, serde_json::json!({"messageId": origin})) + .await; + } + for origin in ["first", "second"] { + server + .send_event( + "assistant.message", + serde_json::json!({ + "messageId":"assistant", "originatingMessageId":origin, "content":origin + }), + ) + .await; + } + server + .send_event("session.idle", serde_json::json!({})) + .await; + for (wait, origin) in waits.into_iter().zip(["first", "second"]) { + let event = timeout(TIMEOUT, wait) + .await + .unwrap() + .unwrap() + .unwrap() + .unwrap(); + assert_eq!(event.data["content"], origin); + } +} + +#[tokio::test] +async fn structured_output_rejects_aborted_missing_and_tool_only_responses() { + for kind in ["missing", "aborted", "tool", "error"] { + let (session, mut server) = create_session_pair().await; + let waiting = tokio::spawn(async move { + session + .send_and_wait( + MessageOptions::new("inventory") + .with_response_schema(serde_json::json!({"type":"object"})), + ) + .await + }); + let request = server.read_request().await; + server + .respond(&request, serde_json::json!({"messageId":"user-1"})) + .await; + server + .send_event( + "user.message", + serde_json::json!({"messageId":"user-1","content":"inventory"}), + ) + .await; + if kind == "tool" { + server.send_event("assistant.message", serde_json::json!({ + "messageId":"assistant", "originatingMessageId":"user-1", "content":"working", + "toolRequests":[{"toolCallId":"tool-1","name":"tool"}] + })).await; + } + if kind == "error" { + server + .send_event( + "session.error", + serde_json::json!({"errorType":"test","message":"provider failed"}), + ) + .await; + } else { + server + .send_event( + "session.idle", + serde_json::json!({"aborted":kind == "aborted"}), + ) + .await; + } + assert!( + timeout(TIMEOUT, waiting).await.unwrap().unwrap().is_err(), + "{kind}" + ); + } +} + +#[tokio::test] +async fn structured_output_timeout_and_cancellation_do_not_block_later_sends() { + let (session, mut server) = create_session_pair().await; + let session = Arc::new(session); + for cancel in [false, true] { + let waiting = tokio::spawn({ + let session = session.clone(); + async move { + session + .send_and_wait( + MessageOptions::new("inventory") + .with_response_schema(serde_json::json!({"type":"object"})) + .with_wait_timeout(Duration::from_millis(50)), + ) + .await + } + }); + let request = server.read_request().await; + server + .respond(&request, serde_json::json!({"messageId":"user-1"})) + .await; + if cancel { + waiting.abort(); + assert!(waiting.await.unwrap_err().is_cancelled()); + } else { + assert!(timeout(TIMEOUT, waiting).await.unwrap().unwrap().is_err()); + } + } + let sending = tokio::spawn(async move { session.send("ordinary").await }); + let request = server.read_request().await; + assert!(request["params"].get("responseFormat").is_none()); + server + .respond(&request, serde_json::json!({"messageId":"plain"})) + .await; + assert_eq!(sending.await.unwrap().unwrap(), "plain"); +} + #[tokio::test] async fn send_and_wait_agent_source_preserves_mode_and_optional_reply() { let (session, mut server) = create_session_pair().await; diff --git a/test/snapshots/structured_output/concurrent_typed_sends_return_their_own_results.yaml b/test/snapshots/structured_output/concurrent_typed_sends_return_their_own_results.yaml new file mode 100644 index 0000000000..24f06c4278 --- /dev/null +++ b/test/snapshots/structured_output/concurrent_typed_sends_return_their_own_results.yaml @@ -0,0 +1,24 @@ +models: + - gpt-4.1 +conversations: + - messages: + - role: system + content: ${system} + - role: user + content: Call first_number exactly once and report its returned number. + - role: assistant + tool_calls: + - id: toolcall_0 + type: function + function: + name: first_number + arguments: "{}" + - role: tool + tool_call_id: toolcall_0 + content: "42" + - role: assistant + content: '{"first":42}' + - role: user + content: What is 30 + 7? Do not use tools. + - role: assistant + content: '{"second":37}' diff --git a/test/snapshots/structured_output/infers_typed_result_after_custom_tool.yaml b/test/snapshots/structured_output/infers_typed_result_after_custom_tool.yaml new file mode 100644 index 0000000000..d3cd70234d --- /dev/null +++ b/test/snapshots/structured_output/infers_typed_result_after_custom_tool.yaml @@ -0,0 +1,24 @@ +models: + - gpt-4.1 +conversations: + - messages: + - role: system + content: ${system} + - role: user + content: Call get_inventory, then report the widget count and color. + - role: assistant + tool_calls: + - id: toolcall_0 + type: function + function: + name: get_inventory + arguments: "{}" + - role: tool + tool_call_id: toolcall_0 + content: The inventory contains 42 red widgets. + - role: assistant + content: '{"color":"red","count":42}' + - role: user + content: Now reply with exactly the plain text HELLO, not JSON. + - role: assistant + content: HELLO diff --git a/test/snapshots/structured_output/send_selects_correlated_response_after_idle.yaml b/test/snapshots/structured_output/send_selects_correlated_response_after_idle.yaml new file mode 100644 index 0000000000..35406a19ba --- /dev/null +++ b/test/snapshots/structured_output/send_selects_correlated_response_after_idle.yaml @@ -0,0 +1,20 @@ +models: + - gpt-4.1 +conversations: + - messages: + - role: system + content: ${system} + - role: user + content: Call read_inventory once, then report the current widget count and color. + - role: assistant + tool_calls: + - id: toolcall_0 + type: function + function: + name: read_inventory + arguments: "{}" + - role: tool + tool_call_id: toolcall_0 + content: The inventory contains 42 red widgets. + - role: assistant + content: '{"color":"red","count":42}' diff --git a/test/snapshots/structured_output/sends_explicit_schema_for_message_and_batch.yaml b/test/snapshots/structured_output/sends_explicit_schema_for_message_and_batch.yaml new file mode 100644 index 0000000000..f450faacb7 --- /dev/null +++ b/test/snapshots/structured_output/sends_explicit_schema_for_message_and_batch.yaml @@ -0,0 +1,16 @@ +models: + - gpt-4.1 +conversations: + - messages: + - role: system + content: ${system} + - role: user + content: There are 42 red widgets in stock. + - role: user + content: Report the widget count and color. + - role: assistant + content: '{"color":"red","count":42}' + - role: user + content: The inventory now has 21 blue widgets. Report the new count and color. + - role: assistant + content: '{"color":"blue","count":21}' diff --git a/test/snapshots/structured_output/typed_result_after_terminal_tool_and_steering.yaml b/test/snapshots/structured_output/typed_result_after_terminal_tool_and_steering.yaml new file mode 100644 index 0000000000..060c0e1a3c --- /dev/null +++ b/test/snapshots/structured_output/typed_result_after_terminal_tool_and_steering.yaml @@ -0,0 +1,22 @@ +models: + - gpt-4.1 +conversations: + - messages: + - role: system + content: ${system} + - role: user + content: Call lookup_number exactly once, then add 5 to the returned number. Do not guess its result. + - role: assistant + tool_calls: + - id: toolcall_0 + type: function + function: + name: lookup_number + arguments: "{}" + - role: tool + tool_call_id: toolcall_0 + content: "58" + - role: user + content: Continue with the original calculation. Do not call any more tools. + - role: assistant + content: '{"answer":63,"contract":"typed_tool"}' diff --git a/test/snapshots/structured_output/typed_wait_returns_late_steering_response.yaml b/test/snapshots/structured_output/typed_wait_returns_late_steering_response.yaml new file mode 100644 index 0000000000..d418eef9f4 --- /dev/null +++ b/test/snapshots/structured_output/typed_wait_returns_late_steering_response.yaml @@ -0,0 +1,14 @@ +models: + - gpt-4.1 +conversations: + - messages: + - role: system + content: ${system} + - role: user + content: What is 19 + 23? Do not use tools. + - role: assistant + content: '{"answer":42}' + - role: user + content: Change the answer to 99. Do not use tools. + - role: assistant + content: '{"answer":99}' diff --git a/test/snapshots/structured_output/typed_wait_returns_stop_hook_correction.yaml b/test/snapshots/structured_output/typed_wait_returns_stop_hook_correction.yaml new file mode 100644 index 0000000000..5f6e6483c4 --- /dev/null +++ b/test/snapshots/structured_output/typed_wait_returns_stop_hook_correction.yaml @@ -0,0 +1,14 @@ +models: + - gpt-4.1 +conversations: + - messages: + - role: system + content: ${system} + - role: user + content: What is 19 + 23? Do not use tools. + - role: assistant + content: '{"answer":42}' + - role: user + content: Correct the answer to 99, not 42. Do not use tools. + - role: assistant + content: '{"answer":99}' diff --git a/test/snapshots/structured_output/typed_wait_returns_stop_hook_correction_after_terminal_tool.yaml b/test/snapshots/structured_output/typed_wait_returns_stop_hook_correction_after_terminal_tool.yaml new file mode 100644 index 0000000000..028d3ee99f --- /dev/null +++ b/test/snapshots/structured_output/typed_wait_returns_stop_hook_correction_after_terminal_tool.yaml @@ -0,0 +1,24 @@ +models: + - gpt-4.1 +conversations: + - messages: + - role: system + content: ${system} + - role: user + content: Call lookup_number exactly once, then add 5 to the returned number. Do not guess its result. + - role: assistant + tool_calls: + - id: toolcall_0 + type: function + function: + name: lookup_number + arguments: "{}" + - role: tool + tool_call_id: toolcall_0 + content: "58" + - role: assistant + content: '{"answer":63}' + - role: user + content: Correct the answer to 99, not 63. Do not use tools. + - role: assistant + content: '{"answer":99}'