diff --git a/README.md b/README.md index da65a1f3d..400ba6eb1 100644 --- a/README.md +++ b/README.md @@ -51,6 +51,8 @@ fx The OpenAI Codex route uses ChatGPT subscription access directly and never sends its OAuth token to Vercel AI Gateway. The session is stored privately at `~/.fx/chatgpt-auth.json` and refreshed when needed. On supported Codex models, `/fast` requests OpenAI's priority service tier and consumes ChatGPT credits at the higher Fast mode rate. +Codex uses HTTP streaming by default. To prefer the WebSocket transport, start fx with `FX_CODEX_TRANSPORT=websocket`. If connection setup fails before the request can be delivered, fx completes that turn over HTTP and keeps using HTTP for the rest of the process. It never replays a request whose delivery is uncertain. WebSocket sessions retain compatible connections and continuation state; concurrent requests use separate ordered lanes rather than sharing one response stream. Each session identity retains at most four lanes by default, and the process retains at most 32 lanes across identities. Set `FX_CODEX_WEBSOCKET_MAX_LANES` or `FX_CODEX_WEBSOCKET_MAX_SLOTS` to a positive integer to choose different limits. + The Grok route uses subscription access directly at xAI and never sends its OAuth token to Vercel AI Gateway or OpenAI. Its session is stored privately at `~/.fx/grok-auth.json`, refreshed when needed, and used only with the authenticated xAI catalog and Responses API. To use an AI Gateway API key instead: diff --git a/src/gateway/codex_websocket_session.zig b/src/gateway/codex_websocket_session.zig new file mode 100644 index 000000000..8a17d2c42 --- /dev/null +++ b/src/gateway/codex_websocket_session.zig @@ -0,0 +1,633 @@ +const std = @import("std"); +const io_mod = @import("../core/shared/io.zig"); +const gateway_client = @import("client.zig"); +const websocket_transport = @import("websocket_transport.zig"); + +const Allocator = std.mem.Allocator; +const pool_alloc = std.heap.c_allocator; + +pub const health_budget: u8 = 3; +pub const default_max_connection_age_ms: i64 = 55 * 60 * 1000; +const default_max_lanes: usize = 4; +const default_max_slots: usize = 32; +const max_connection_age_env = "FX_CODEX_WEBSOCKET_MAX_CONNECTION_AGE_MS"; +const max_lanes_env = "FX_CODEX_WEBSOCKET_MAX_LANES"; +const max_slots_env = "FX_CODEX_WEBSOCKET_MAX_SLOTS"; + +const Slot = struct { + session_id: []u8, + account_id: []u8, + model: []u8, + endpoint: []u8, + authorization_fingerprint: [std.crypto.hash.sha2.Sha256.digest_length]u8, + connection: ?*websocket_transport.Connection, + busy: bool, + health_failures: u8, + opened_at_ms: i64, + last_used_at_ms: i64, + continuation_response_id: ?[]u8, + continuation_baseline: ?[]u8, + continuation_durable_baseline: ?[]u8, + continuation_shape: [std.crypto.hash.sha2.Sha256.digest_length]u8, + continuation_valid: bool, + + fn clearContinuation(self: *Slot) void { + if (self.continuation_response_id) |value| pool_alloc.free(value); + if (self.continuation_baseline) |value| pool_alloc.free(value); + if (self.continuation_durable_baseline) |value| pool_alloc.free(value); + self.continuation_response_id = null; + self.continuation_baseline = null; + self.continuation_durable_baseline = null; + self.continuation_valid = false; + } + + fn deinit(self: *Slot) void { + if (self.connection) |connection| websocket_transport.close(connection, pool_alloc); + self.clearContinuation(); + pool_alloc.free(self.session_id); + pool_alloc.free(self.account_id); + pool_alloc.free(self.model); + pool_alloc.free(self.endpoint); + self.* = undefined; + } +}; + +var pool_mutex: std.Io.Mutex = .init; +var slots: std.ArrayList(Slot) = .empty; + +pub const AcquireArgs = struct { + session_id: ?[]const u8, + account_id: []const u8, + model: []const u8, + endpoint: []const u8, + authorization: []const u8, + deadline: ?std.Io.Clock.Timestamp, + cancel_flag: *std.atomic.Value(bool), + delivery: *gateway_client.DeliveryCertainty, + continuation_input: ?[]const u8 = null, + continuation_shape: ?[std.crypto.hash.sha2.Sha256.digest_length]u8 = null, +}; + +pub const Checkout = struct { + slot: ?usize, + connection: *websocket_transport.Connection, + reused: bool, + retained: bool, + handshake_ms: i64, + health_failures: u8, +}; + +pub const Continuation = struct { + previous_response_id: []const u8, + delta_input: []const u8, +}; + +fn continuationDelta(full_input: []const u8, baseline: []const u8) ?[]const u8 { + if (!std.mem.startsWith(u8, full_input, baseline)) return null; + if (full_input.len == baseline.len) return ""; + if (full_input[baseline.len] != ',') return null; + return full_input[baseline.len + 1 ..]; +} + +pub const Outcome = enum { completed, failed }; + +fn sessionKey(session_id: ?[]const u8) []const u8 { + const value = session_id orelse return ""; + return if (value.len == 0) "" else value; +} + +fn authorizationFingerprint(authorization: []const u8) [std.crypto.hash.sha2.Sha256.digest_length]u8 { + var digest: [std.crypto.hash.sha2.Sha256.digest_length]u8 = undefined; + std.crypto.hash.sha2.Sha256.hash(authorization, &digest, .{}); + return digest; +} + +fn matches(slot: *const Slot, args: AcquireArgs) bool { + const fingerprint = authorizationFingerprint(args.authorization); + return std.mem.eql(u8, slot.session_id, sessionKey(args.session_id)) and + std.mem.eql(u8, slot.account_id, args.account_id) and + std.mem.eql(u8, slot.model, args.model) and + std.mem.eql(u8, slot.endpoint, args.endpoint) and + std.mem.eql(u8, &slot.authorization_fingerprint, &fingerprint); +} + +fn maxConnectionAgeMs() !i64 { + const value = io_mod.getenv(max_connection_age_env) orelse return default_max_connection_age_ms; + const parsed = std.fmt.parseInt(i64, value, 10) catch return error.InvalidOpenAICodexTransport; + if (parsed < 0) return error.InvalidOpenAICodexTransport; + return parsed; +} + +fn maxLanes() !usize { + const value = io_mod.getenv(max_lanes_env) orelse return default_max_lanes; + const parsed = std.fmt.parseInt(usize, value, 10) catch return error.InvalidOpenAICodexTransport; + if (parsed == 0) return error.InvalidOpenAICodexTransport; + return parsed; +} + +fn maxSlots() !usize { + const value = io_mod.getenv(max_slots_env) orelse return default_max_slots; + const parsed = std.fmt.parseInt(usize, value, 10) catch return error.InvalidOpenAICodexTransport; + if (parsed == 0) return error.InvalidOpenAICodexTransport; + return parsed; +} + +fn initSlot(args: AcquireArgs, busy: bool) !Slot { + const session_id = try pool_alloc.dupe(u8, sessionKey(args.session_id)); + errdefer pool_alloc.free(session_id); + const account_id = try pool_alloc.dupe(u8, args.account_id); + errdefer pool_alloc.free(account_id); + const model = try pool_alloc.dupe(u8, args.model); + errdefer pool_alloc.free(model); + const endpoint = try pool_alloc.dupe(u8, args.endpoint); + errdefer pool_alloc.free(endpoint); + return .{ + .session_id = session_id, + .account_id = account_id, + .model = model, + .endpoint = endpoint, + .authorization_fingerprint = authorizationFingerprint(args.authorization), + .connection = null, + .busy = busy, + .health_failures = 0, + .opened_at_ms = 0, + .last_used_at_ms = io_mod.milliTimestamp(), + .continuation_response_id = null, + .continuation_baseline = null, + .continuation_durable_baseline = null, + .continuation_shape = undefined, + .continuation_valid = false, + }; +} + +fn appendSlot(args: AcquireArgs, busy: bool) !usize { + try slots.append(pool_alloc, try initSlot(args, busy)); + return slots.items.len - 1; +} + +fn continuationMatches(slot: *const Slot, full_input: []const u8, shape: [std.crypto.hash.sha2.Sha256.digest_length]u8) bool { + if (!slot.continuation_valid or !std.mem.eql(u8, &slot.continuation_shape, &shape)) return false; + if (slot.continuation_baseline) |baseline| { + if (continuationDelta(full_input, baseline) != null) return true; + } + if (slot.continuation_durable_baseline) |baseline| { + if (continuationDelta(full_input, baseline) != null) return true; + } + return false; +} + +const LaneSelection = struct { + index: ?usize, + matching_count: usize, +}; + +fn selectIdleLane(slot_items: []Slot, args: AcquireArgs) LaneSelection { + var first_idle: ?usize = null; + var continuation_idle: ?usize = null; + var matching_count: usize = 0; + for (slot_items, 0..) |*slot, index| { + if (!matches(slot, args)) continue; + matching_count += 1; + if (slot.busy) continue; + if (first_idle == null) first_idle = index; + if (args.continuation_input) |full_input| { + if (args.continuation_shape) |shape| { + if (continuationMatches(slot, full_input, shape)) { + continuation_idle = index; + break; + } + } + } + } + return .{ .index = continuation_idle orelse first_idle, .matching_count = matching_count }; +} + +const LaneChoice = union(enum) { + existing: usize, + append, + temporary, +}; + +fn chooseLane(slot_items: []Slot, args: AcquireArgs, lane_limit: usize) LaneChoice { + const selection = selectIdleLane(slot_items, args); + if (selection.index) |index| return .{ .existing = index }; + if (selection.matching_count < lane_limit) return .append; + return .temporary; +} + +fn incrementFailure(slot: *Slot) void { + slot.health_failures = std.math.add(u8, slot.health_failures, 1) catch std.math.maxInt(u8); +} + +fn leastRecentlyUsedIdle(slot_items: []Slot) ?usize { + var selected: ?usize = null; + for (slot_items, 0..) |*slot, index| { + if (slot.busy) continue; + if (selected == null or slot.last_used_at_ms < slot_items[selected.?].last_used_at_ms) { + selected = index; + } + } + return selected; +} + +fn incompatibleIdle(slot_items: []Slot, args: AcquireArgs) ?usize { + for (slot_items, 0..) |*slot, index| { + if (slot.busy or matches(slot, args)) continue; + if (std.mem.eql(u8, slot.session_id, sessionKey(args.session_id))) return index; + } + return null; +} + +fn replaceSlot(index: usize, args: AcquireArgs) !?*websocket_transport.Connection { + const replacement = try initSlot(args, true); + const displaced = slots.items[index].connection; + slots.items[index].connection = null; + slots.items[index].deinit(); + slots.items[index] = replacement; + return displaced; +} + +pub fn acquire(_: Allocator, args: AcquireArgs) !Checkout { + if (args.cancel_flag.load(.seq_cst)) return error.Cancelled; + if (args.deadline) |deadline| { + const now = std.Io.Clock.Timestamp.now(io_mod.getIo(), .awake); + if (!std.Io.Clock.Timestamp.compare(now, .lt, deadline)) return error.Timeout; + } + + const lane_limit = try maxLanes(); + const slot_limit = try maxSlots(); + const age_limit = try maxConnectionAgeMs(); + var index: ?usize = null; + var retained = true; + var reusable: ?*websocket_transport.Connection = null; + var displaced: ?*websocket_transport.Connection = null; + var prior_health: u8 = 0; + + pool_mutex.lockUncancelable(io_mod.getIo()); + switch (chooseLane(slots.items, args, lane_limit)) { + .existing => |existing| { + index = existing; + const slot = &slots.items[existing]; + slot.busy = true; + prior_health = slot.health_failures; + const expired = age_limit != 0 and + io_mod.milliTimestamp() - slot.opened_at_ms > age_limit; + if (slot.connection != null and slot.health_failures < health_budget and !expired) { + reusable = slot.connection; + } else { + displaced = slot.connection; + slot.connection = null; + slot.clearContinuation(); + } + }, + .append => { + if (incompatibleIdle(slots.items, args)) |victim| { + index = victim; + displaced = replaceSlot(victim, args) catch |err| { + pool_mutex.unlock(io_mod.getIo()); + return err; + }; + } else if (slots.items.len < slot_limit) { + index = appendSlot(args, true) catch |err| { + pool_mutex.unlock(io_mod.getIo()); + return err; + }; + } else if (leastRecentlyUsedIdle(slots.items)) |victim| { + index = victim; + displaced = replaceSlot(victim, args) catch |err| { + pool_mutex.unlock(io_mod.getIo()); + return err; + }; + } else { + retained = false; + } + }, + .temporary => retained = false, + } + pool_mutex.unlock(io_mod.getIo()); + + // Socket close, health checks, and connection establishment are all + // deliberately outside the global pool mutex. + if (displaced) |connection| websocket_transport.close(connection, pool_alloc); + if (reusable) |connection| { + websocket_transport.ping(connection, args.cancel_flag, args.deadline, args.delivery) catch |err| { + websocket_transport.close(connection, pool_alloc); + pool_mutex.lockUncancelable(io_mod.getIo()); + if (index) |slot_index| { + const slot = &slots.items[slot_index]; + if (slot.connection == connection) slot.connection = null; + slot.clearContinuation(); + incrementFailure(slot); + prior_health = slot.health_failures; + } + pool_mutex.unlock(io_mod.getIo()); + if (err == error.Cancelled) { + rollbackReservation(index); + return err; + } + reusable = null; + }; + if (reusable != null) { + return .{ + .slot = index, + .connection = connection, + .reused = true, + .retained = retained, + .handshake_ms = 0, + .health_failures = prior_health, + }; + } + } + + const started_at_ms = io_mod.milliTimestamp(); + const connection = websocket_transport.connect(pool_alloc, .{ + .endpoint = args.endpoint, + .authorization = args.authorization, + .account_id = args.account_id, + .session_id = args.session_id, + .deadline = args.deadline, + .cancel_flag = args.cancel_flag, + .delivery = args.delivery, + }) catch |err| { + rollbackReservation(index); + return err; + }; + if (index) |slot_index| { + pool_mutex.lockUncancelable(io_mod.getIo()); + const slot = &slots.items[slot_index]; + slot.connection = connection; + slot.clearContinuation(); + slot.opened_at_ms = connection.opened_at_ms; + slot.last_used_at_ms = io_mod.milliTimestamp(); + pool_mutex.unlock(io_mod.getIo()); + } + return .{ + .slot = index, + .connection = connection, + .reused = false, + .retained = retained, + .handshake_ms = @max(io_mod.milliTimestamp() - started_at_ms, 0), + .health_failures = prior_health, + }; +} + +fn rollbackReservation(index: ?usize) void { + const slot_index = index orelse return; + pool_mutex.lockUncancelable(io_mod.getIo()); + if (slot_index < slots.items.len) slots.items[slot_index].busy = false; + pool_mutex.unlock(io_mod.getIo()); +} + +pub fn continuation( + index: ?usize, + full_input: []const u8, + shape: [std.crypto.hash.sha2.Sha256.digest_length]u8, +) ?Continuation { + const slot_index = index orelse return null; + pool_mutex.lockUncancelable(io_mod.getIo()); + defer pool_mutex.unlock(io_mod.getIo()); + if (slot_index >= slots.items.len) return null; + const slot = &slots.items[slot_index]; + if (!slot.busy or !slot.continuation_valid) return null; + if (!std.mem.eql(u8, &slot.continuation_shape, &shape)) { + slot.clearContinuation(); + return null; + } + const response_id = slot.continuation_response_id orelse return null; + const baseline = slot.continuation_baseline orelse return null; + const delta = continuationDelta(full_input, baseline) orelse durable: { + const durable_baseline = slot.continuation_durable_baseline orelse { + slot.clearContinuation(); + return null; + }; + break :durable continuationDelta(full_input, durable_baseline) orelse { + slot.clearContinuation(); + return null; + }; + }; + return .{ + .previous_response_id = response_id, + .delta_input = delta, + }; +} + +pub fn recordCompletion( + index: ?usize, + response_id: []const u8, + baseline: []const u8, + durable_baseline: []const u8, + shape: [std.crypto.hash.sha2.Sha256.digest_length]u8, +) void { + const slot_index = index orelse return; + pool_mutex.lockUncancelable(io_mod.getIo()); + defer pool_mutex.unlock(io_mod.getIo()); + if (slot_index >= slots.items.len) return; + const slot = &slots.items[slot_index]; + slot.clearContinuation(); + const owned_id = pool_alloc.dupe(u8, response_id) catch return; + const owned_baseline = pool_alloc.dupe(u8, baseline) catch { + pool_alloc.free(owned_id); + return; + }; + const owned_durable_baseline = pool_alloc.dupe(u8, durable_baseline) catch { + pool_alloc.free(owned_id); + pool_alloc.free(owned_baseline); + return; + }; + slot.continuation_response_id = owned_id; + slot.continuation_baseline = owned_baseline; + slot.continuation_durable_baseline = owned_durable_baseline; + slot.continuation_shape = shape; + slot.continuation_valid = true; +} + +pub fn release(checkout: Checkout, outcome: Outcome) void { + const index = checkout.slot orelse { + websocket_transport.close(checkout.connection, pool_alloc); + return; + }; + var discarded: ?*websocket_transport.Connection = null; + pool_mutex.lockUncancelable(io_mod.getIo()); + if (index >= slots.items.len) { + pool_mutex.unlock(io_mod.getIo()); + websocket_transport.close(checkout.connection, pool_alloc); + return; + } + const slot = &slots.items[index]; + slot.busy = false; + slot.last_used_at_ms = io_mod.milliTimestamp(); + switch (outcome) { + .completed => slot.health_failures = 0, + .failed => { + incrementFailure(slot); + discarded = slot.connection; + slot.connection = null; + slot.clearContinuation(); + }, + } + pool_mutex.unlock(io_mod.getIo()); + if (discarded) |connection| websocket_transport.close(connection, pool_alloc); +} + +pub fn shutdown() void { + pool_mutex.lockUncancelable(io_mod.getIo()); + const retired = slots; + slots = .empty; + pool_mutex.unlock(io_mod.getIo()); + var owned = retired; + for (owned.items) |*slot| slot.deinit(); + owned.deinit(pool_alloc); +} + +test "retained Codex WebSocket identity includes authorization without storing it" { + const slot = Slot{ + .session_id = @constCast("session-a"), + .account_id = @constCast("account-a"), + .model = @constCast("gpt-5.6-sol"), + .endpoint = @constCast("http://127.0.0.1/responses"), + .authorization_fingerprint = authorizationFingerprint("Bearer token-a"), + .connection = null, + .busy = false, + .health_failures = 0, + .opened_at_ms = 0, + .last_used_at_ms = 0, + .continuation_response_id = null, + .continuation_baseline = null, + .continuation_durable_baseline = null, + .continuation_shape = undefined, + .continuation_valid = false, + }; + const base = AcquireArgs{ + .session_id = "session-a", + .account_id = "account-a", + .model = "gpt-5.6-sol", + .endpoint = "http://127.0.0.1/responses", + .authorization = "Bearer token-a", + .deadline = null, + .cancel_flag = undefined, + .delivery = undefined, + }; + + try std.testing.expect(matches(&slot, base)); + + var rotated = base; + rotated.authorization = "Bearer token-b"; + try std.testing.expect(!matches(&slot, rotated)); + + var changed_model = base; + changed_model.model = "gpt-5.4"; + try std.testing.expect(!matches(&slot, changed_model)); +} + +test "Codex WebSocket continuation requires an exact item boundary prefix" { + try std.testing.expectEqualStrings( + "{\"role\":\"user\",\"content\":[]}", + continuationDelta( + "{\"type\":\"message\"},{\"role\":\"user\",\"content\":[]}", + "{\"type\":\"message\"}", + ).?, + ); + try std.testing.expectEqualStrings( + "", + continuationDelta("{\"type\":\"message\"}", "{\"type\":\"message\"}").?, + ); + try std.testing.expect(continuationDelta("{\"type\":\"message\"}suffix", "{\"type\":\"message\"}") == null); + try std.testing.expect(continuationDelta("{\"type\":\"other\"}", "{\"type\":\"message\"}") == null); +} + +test "Codex WebSocket lane selection preserves continuation affinity" { + shutdown(); + defer shutdown(); + + var cancel_flag = std.atomic.Value(bool).init(false); + var delivery = gateway_client.DeliveryCertainty.init(); + var shape: [std.crypto.hash.sha2.Sha256.digest_length]u8 = undefined; + std.crypto.hash.sha2.Sha256.hash("shape", &shape, .{}); + const args = AcquireArgs{ + .session_id = "session-a", + .account_id = "account-a", + .model = "gpt-5.6-sol", + .endpoint = "http://127.0.0.1/responses", + .authorization = "Bearer token-a", + .deadline = null, + .cancel_flag = &cancel_flag, + .delivery = &delivery, + .continuation_input = "{\"type\":\"message\"},{\"role\":\"user\"}", + .continuation_shape = shape, + }; + + const first = try appendSlot(args, false); + slots.items[first].busy = true; + const second = try appendSlot(args, false); + slots.items[second].busy = true; + recordCompletion(second, "response-2", "{\"type\":\"message\"}", "{\"type\":\"message\"}", shape); + slots.items[second].busy = false; + + const selection = chooseLane(slots.items, args, 2); + try std.testing.expectEqual(second, selection.existing); + + slots.items[second].busy = true; + try std.testing.expect(chooseLane(slots.items, args, 2) == .temporary); + try std.testing.expect(chooseLane(slots.items, args, 3) == .append); +} + +test "Codex WebSocket global slot eviction selects the least recently used idle lane" { + shutdown(); + defer shutdown(); + + var cancel_flag = std.atomic.Value(bool).init(false); + var delivery = gateway_client.DeliveryCertainty.init(); + const args = AcquireArgs{ + .session_id = "session-a", + .account_id = "account-a", + .model = "gpt-5.6-sol", + .endpoint = "http://127.0.0.1/responses", + .authorization = "Bearer token-a", + .deadline = null, + .cancel_flag = &cancel_flag, + .delivery = &delivery, + }; + _ = try appendSlot(args, false); + _ = try appendSlot(args, false); + _ = try appendSlot(args, false); + slots.items[0].last_used_at_ms = 30; + slots.items[1].last_used_at_ms = 10; + slots.items[2].last_used_at_ms = 20; + slots.items[1].busy = true; + + try std.testing.expectEqual(@as(?usize, 2), leastRecentlyUsedIdle(slots.items)); + try std.testing.expectEqual(@as(usize, 3), slots.items.len); +} + +test "Codex WebSocket slot storage remains bounded under identity churn" { + shutdown(); + defer shutdown(); + + var cancel_flag = std.atomic.Value(bool).init(false); + var delivery = gateway_client.DeliveryCertainty.init(); + var session_buffer: [32]u8 = undefined; + var args = AcquireArgs{ + .session_id = "", + .account_id = "account-a", + .model = "gpt-5.6-sol", + .endpoint = "http://127.0.0.1/responses", + .authorization = "Bearer token-a", + .deadline = null, + .cancel_flag = &cancel_flag, + .delivery = &delivery, + }; + const limit: usize = 3; + for (0..64) |identity| { + args.session_id = try std.fmt.bufPrint(&session_buffer, "session-{d}", .{identity}); + if (slots.items.len < limit) { + _ = try appendSlot(args, false); + } else { + const victim = leastRecentlyUsedIdle(slots.items).?; + _ = try replaceSlot(victim, args); + slots.items[victim].busy = false; + slots.items[victim].last_used_at_ms = @intCast(identity); + } + try std.testing.expect(slots.items.len <= limit); + } + try std.testing.expectEqual(limit, slots.items.len); +} diff --git a/src/gateway/openai_codex.zig b/src/gateway/openai_codex.zig index 95e3356ed..fedf1e211 100644 --- a/src/gateway/openai_codex.zig +++ b/src/gateway/openai_codex.zig @@ -7,6 +7,8 @@ const io_mod = @import("../core/shared/io.zig"); const types = @import("../core/shared/types.zig"); const gateway_client = @import("client.zig"); const responses_protocol = @import("responses_protocol.zig"); +const codex_websocket = @import("openai_codex_websocket.zig"); +const debug_trace = @import("../core/shared/debug_trace.zig"); const model_tool_schema = @import("../core/tooling/model_tool_schema.zig"); const Allocator = std.mem.Allocator; @@ -22,6 +24,41 @@ const max_tool_arguments_bytes: usize = 4 * 1024 * 1024; const max_provider_state_bytes: usize = 4 * 1024 * 1024; const transfer_buffer_bytes: usize = 256 * 1024; const connect_timeout_ms: i64 = 30_000; +const transport_env = "FX_CODEX_TRANSPORT"; + +const Transport = enum { sse, websocket }; +var sse_fallback_active = std.atomic.Value(bool).init(false); + +fn selectedTransport() !Transport { + const value = io_mod.getenv(transport_env) orelse return .sse; + if (std.mem.eql(u8, value, "sse") or std.mem.eql(u8, value, "auto")) return .sse; + if (std.mem.eql(u8, value, "websocket")) { + return if (sse_fallback_active.load(.seq_cst)) .sse else .websocket; + } + return error.InvalidOpenAICodexTransport; +} + +fn allowsSseFallback(err: anyerror, delivery: gateway_client.DeliveryCertainty.State) bool { + if (delivery != .definitely_unsent) return false; + return switch (err) { + error.Cancelled, + error.OutOfMemory, + error.ProviderAdmissionMissing, + error.ProviderAdmissionRepeated, + error.InvalidOpenAICodexTransport, + => false, + else => true, + }; +} + +fn armSseFallback(err: anyerror) void { + sse_fallback_active.store(true, .seq_cst); + debug_trace.logf( + "stream", + "Codex WebSocket transport disabled for this process error={s}", + .{@errorName(err)}, + ); +} const CodexLimits = struct { aggregate_bytes: usize = max_sse_aggregate_bytes, @@ -32,10 +69,25 @@ const CodexLimits = struct { provider_state_bytes: usize = max_provider_state_bytes, }; +fn codexStreamLimits(limits: CodexLimits) responses_protocol.StreamLimits { + return .{ + .aggregate_bytes = limits.aggregate_bytes, + .events = limits.events, + .tool_calls = limits.tool_calls, + .tool_identity_bytes = limits.tool_identity_bytes, + .tool_arguments_bytes = limits.tool_arguments_bytes, + .provider_state_bytes = limits.provider_state_bytes, + }; +} + pub const agent_stream_provider = stream_provider.Provider{ .stream_fn = streamCompletion, }; +pub fn shutdownWebSockets() void { + codex_websocket.shutdown(); +} + fn validateModel(model: []const u8) !void { if (model.len == 0 or model.len > 1024) return error.InvalidOpenAICodexModel; for (model) |byte| { @@ -47,12 +99,75 @@ pub fn buildRequest( alloc: Allocator, request: stream_provider.RequestData, ) ![]u8 { + var out: std.Io.Writer.Allocating = .init(alloc); + errdefer out.deinit(); + try writeResponseRequestStart(&out.writer, alloc, request, null); + try out.writer.writeAll(",\"store\":false,\"stream\":true"); + try out.writer.writeByte('}'); + return out.toOwnedSlice(); +} + +fn buildWebSocketRequest( + alloc: Allocator, + request: stream_provider.RequestData, +) ![]u8 { + const input = try buildResponseInput(alloc, request.messages, request.verified_images); + defer alloc.free(input); + return buildWebSocketRequestWithInput(alloc, request, input, null); +} + +fn buildWebSocketRequestWithInput( + alloc: Allocator, + request: stream_provider.RequestData, + input: []const u8, + previous_response_id: ?[]const u8, +) ![]u8 { + var out: std.Io.Writer.Allocating = .init(alloc); + errdefer out.deinit(); + try writeResponseRequestStartWithInput(&out.writer, alloc, request, "response.create", input); + if (previous_response_id) |response_id| { + try out.writer.writeAll(",\"previous_response_id\":"); + try std.json.Stringify.value(response_id, .{}, &out.writer); + } + try out.writer.writeByte('}'); + return out.toOwnedSlice(); +} + +/// Writes the fields shared by SSE and WebSocket Responses envelopes. The +/// caller owns the opening and closing JSON object delimiters. +fn writeResponseRequestStart( + writer: *std.Io.Writer, + alloc: Allocator, + request: stream_provider.RequestData, + websocket_type: ?[]const u8, +) !void { + const input = try buildResponseInput(alloc, request.messages, request.verified_images); + defer alloc.free(input); + return writeResponseRequestStartWithInput(writer, alloc, request, websocket_type, input); +} + +fn writeResponseRequestStartWithInput( + writer: *std.Io.Writer, + alloc: Allocator, + request: stream_provider.RequestData, + websocket_type: ?[]const u8, + input: []const u8, +) !void { try validateModel(request.model); if (request.budget) |budget| { if (budget.cancel_flag) |flag| if (flag.load(.seq_cst)) return error.Cancelled; _ = budget.deadline; } + try writer.writeByte('{'); + if (websocket_type) |value| { + try writer.writeAll("\"type\":"); + try std.json.Stringify.value(value, .{}, writer); + try writer.writeByte(','); + } + try writer.writeAll("\"model\":"); + try std.json.Stringify.value(request.model, .{}, writer); + var instructions: std.Io.Writer.Allocating = .init(alloc); defer instructions.deinit(); for (request.messages) |message| { @@ -64,15 +179,10 @@ pub fn buildRequest( } if (instructions.written().len == 0) try instructions.writer.writeAll("You are a helpful assistant."); - var out: std.Io.Writer.Allocating = .init(alloc); - errdefer out.deinit(); - const writer = &out.writer; - try writer.writeAll("{\"model\":"); - try std.json.Stringify.value(request.model, .{}, writer); - try writer.writeAll(",\"store\":false,\"stream\":true,\"instructions\":"); + try writer.writeAll(",\"instructions\":"); try std.json.Stringify.value(instructions.written(), .{}, writer); try writer.writeAll(",\"input\":["); - try writeResponsesInput(writer, alloc, request.messages, request.verified_images); + try writer.writeAll(input); try writer.writeByte(']'); _ = try responses_protocol.writeTools(writer, alloc, request.tools); @@ -104,10 +214,29 @@ pub fn buildRequest( } // The ChatGPT Codex endpoint chooses the model's output limit and rejects // the public Responses API max_output_tokens parameter. - try writer.writeByte('}'); +} + +fn buildResponseInput( + alloc: Allocator, + messages: []const types.ChatMessage, + images: ?[]const image_attachments.VerifiedSnapshot, +) ![]u8 { + var out: std.Io.Writer.Allocating = .init(alloc); + errdefer out.deinit(); + try writeResponsesInput(&out.writer, alloc, messages, images); return out.toOwnedSlice(); } +fn buildDurableResponseInput( + alloc: Allocator, + messages: []const types.ChatMessage, +) ![]u8 { + const projected = try alloc.dupe(types.ChatMessage, messages); + defer alloc.free(projected); + for (projected) |*message| message.provider_state_json = null; + return buildResponseInput(alloc, projected, null); +} + fn writeResponsesInput( writer: *std.Io.Writer, alloc: Allocator, @@ -138,15 +267,37 @@ fn streamCompletion( return stream_provider.failResult(error.CodexSubscriptionCredentialRequired); } try validateModel(request.model); + const transport = try selectedTransport(); + if (transport == .websocket) { + if (streamWebSocketPrepared(alloc, request)) |result| return result else |err| { + if (request.cancel_flag.load(.seq_cst)) return stream_provider.failResult(error.Cancelled); + const delivery = request.delivery.load(); + if (!allowsSseFallback(err, delivery)) { + request.attempt_evidence.network_failure = gateway_client.networkFailureEvidence(err, delivery); + return err; + } + armSseFallback(err); + } + } + const payload = try buildRequest(alloc, request.data()); defer alloc.free(payload); - return streamPrepared(alloc, request, payload) catch |err| { + return streamPreparedWithAdmission( + alloc, + request, + payload, + !request.attempt_evidence.provider_admitted, + ) catch |err| { if (request.cancel_flag.load(.seq_cst)) return stream_provider.failResult(error.Cancelled); request.attempt_evidence.network_failure = gateway_client.networkFailureEvidence(err, request.delivery.load()); return err; }; } +fn admitCodexTransport(admission: stream_provider.Admission) !void { + try admission.admit(); +} + const OpenedRequest = struct { request: ?std.http.Client.Request, @@ -187,23 +338,24 @@ pub fn streamPrepared( alloc: Allocator, request: stream_provider.ModelRequest, payload: []const u8, +) !stream_provider.Result { + return streamPreparedWithAdmission(alloc, request, payload, true); +} + +fn streamPreparedWithAdmission( + alloc: Allocator, + request: stream_provider.ModelRequest, + payload: []const u8, + should_admit: bool, ) !stream_provider.Result { if (request.cancel_flag.load(.seq_cst)) return stream_provider.failResult(error.Cancelled); - const account_id = try chatgpt_oauth.extractAccountId(alloc, request.credential.secret); - defer alloc.free(account_id); - const auth_header = try std.fmt.allocPrint(alloc, "Bearer {s}", .{request.credential.secret}); - defer secret.zeroAndFree(alloc, auth_header); - const request_endpoint = if (io_mod.getenv(e2e_endpoint_env)) |override| endpoint: { - if (!gateway_client.isLoopbackHttpUrl(override)) { - return stream_provider.failResult(error.InvalidE2EOpenAICodexEndpoint); - } - break :endpoint override; - } else endpoint; - const uri = try std.Uri.parse(request_endpoint); + var prepared = try prepareCodexTransport(alloc, request); + defer prepared.deinit(alloc); + const uri = try std.Uri.parse(prepared.endpoint); var extra_headers_buf: [7]std.http.Header = undefined; var extra_count: usize = 0; - extra_headers_buf[extra_count] = .{ .name = "chatgpt-account-id", .value = account_id }; + extra_headers_buf[extra_count] = .{ .name = "chatgpt-account-id", .value = prepared.account_id }; extra_count += 1; extra_headers_buf[extra_count] = .{ .name = "originator", .value = "fx" }; extra_count += 1; @@ -223,14 +375,14 @@ pub fn streamPrepared( var open_operation = OpenRequestOperation{ .client = &client, .uri = uri, - .auth_header = auth_header, + .auth_header = prepared.authorization, .extra_headers = extra_headers_buf[0..extra_count], }; const connect_deadline = std.Io.Clock.Timestamp.fromNow(io_mod.getIo(), .{ .clock = .awake, .raw = .fromMilliseconds(connect_timeout_ms), }); - try request.admission.admit(); + if (should_admit) try admitCodexTransport(request.admission); var opened = try gateway_client.runBoundedHttpOperation( OpenedRequest, alloc, @@ -321,6 +473,61 @@ pub fn streamPrepared( } }; } +const PreparedCodexTransport = struct { + account_id: []u8, + authorization: []u8, + endpoint: []const u8, + + fn deinit(self: *PreparedCodexTransport, alloc: Allocator) void { + alloc.free(self.account_id); + secret.zeroAndFree(alloc, self.authorization); + self.* = undefined; + } +}; + +fn prepareCodexTransport( + alloc: Allocator, + request: stream_provider.ModelRequest, +) !PreparedCodexTransport { + const account_id = try chatgpt_oauth.extractAccountId(alloc, request.credential.secret); + errdefer alloc.free(account_id); + const authorization = try std.fmt.allocPrint(alloc, "Bearer {s}", .{request.credential.secret}); + errdefer secret.zeroAndFree(alloc, authorization); + const request_endpoint = if (io_mod.getenv(e2e_endpoint_env)) |override| endpoint: { + if (!gateway_client.isLoopbackHttpUrl(override)) return error.InvalidE2EOpenAICodexEndpoint; + break :endpoint override; + } else endpoint; + return .{ + .account_id = account_id, + .authorization = authorization, + .endpoint = request_endpoint, + }; +} + +fn streamWebSocketPrepared( + alloc: Allocator, + request: stream_provider.ModelRequest, +) !stream_provider.Result { + if (request.cancel_flag.load(.seq_cst)) return error.Cancelled; + var prepared = try prepareCodexTransport(alloc, request); + defer prepared.deinit(alloc); + return codex_websocket.stream( + alloc, + request, + .{ + .endpoint = prepared.endpoint, + .authorization = prepared.authorization, + .account_id = prepared.account_id, + }, + .{ + .build_input = buildResponseInput, + .build_request = buildWebSocketRequestWithInput, + .build_durable_input = buildDurableResponseInput, + }, + codexStreamLimits(.{}), + ); +} + const EventBridge = struct { fn sink(raw: *anyopaque) *stream_provider.EventSink { return @ptrCast(@alignCast(raw)); @@ -438,14 +645,7 @@ fn consumeSse( .on_reasoning = on_reasoning_chunk, .on_tool_input = on_tool_input_chunk, }; - const stream_limits = responses_protocol.StreamLimits{ - .aggregate_bytes = limits.aggregate_bytes, - .events = limits.events, - .tool_calls = limits.tool_calls, - .tool_identity_bytes = limits.tool_identity_bytes, - .tool_arguments_bytes = limits.tool_arguments_bytes, - .provider_state_bytes = limits.provider_state_bytes, - }; + const stream_limits = codexStreamLimits(limits); while (try sse.next(alloc, reader)) |json_text| { defer sse.release(); if (reducer.applyJson( @@ -464,6 +664,7 @@ fn consumeSse( fn mapReducerError(err: anyerror) anyerror { return switch (err) { error.InvalidEvent => error.InvalidOpenAICodexSseEvent, + error.PreviousResponseNotFound => error.PreviousResponseNotFound, error.ResponseFailed => error.OpenAICodexResponseFailed, error.StreamIncomplete => error.OpenAICodexStreamIncomplete, error.ToolCallLimitExceeded => error.OpenAICodexToolCallLimitExceeded, @@ -473,6 +674,59 @@ fn mapReducerError(err: anyerror) anyerror { }; } +test "OpenAI Codex WebSocket request uses the shared response request fields" { + const messages = [_]types.ChatMessage{ + .{ .role = .system, .content = "Be concise." }, + .{ .role = .user, .content = "Read it." }, + }; + const request: stream_provider.RequestData = .{ + .model = "gpt-5.4", + .messages = &messages, + .tool_choice = .auto, + .provider_options = .{}, + }; + const sse_payload = try buildRequest(std.testing.allocator, request); + defer std.testing.allocator.free(sse_payload); + const websocket_payload = try buildWebSocketRequest(std.testing.allocator, request); + defer std.testing.allocator.free(websocket_payload); + + try std.testing.expect(std.mem.find(u8, websocket_payload, "\"type\":\"response.create\"") != null); + try std.testing.expect(std.mem.find(u8, websocket_payload, "\"model\":\"gpt-5.4\"") != null); + try std.testing.expect(std.mem.find(u8, websocket_payload, "\"instructions\":\"Be concise.\"") != null); + try std.testing.expect(std.mem.find(u8, websocket_payload, "\"stream\"") == null); + try std.testing.expect(std.mem.find(u8, websocket_payload, "\"store\"") == null); + try std.testing.expect(std.mem.find(u8, sse_payload, "\"stream\":true") != null); + try std.testing.expect(std.mem.find(u8, sse_payload, "\"store\":false") != null); +} + +test "OpenAI Codex transport admission invokes the shared admission boundary" { + const Capture = struct { + called: bool = false, + + fn admit(raw: *anyopaque) !void { + const self: *@This() = @ptrCast(@alignCast(raw)); + self.called = true; + } + }; + var capture: Capture = .{}; + try admitCodexTransport(.{ .context = &capture, .admit_fn = Capture.admit }); + try std.testing.expect(capture.called); +} + +test "OpenAI Codex transport policy keeps auto on SSE" { + // Environment-dependent selection is covered by integration launch tests. + if (io_mod.getenv(transport_env) == null) { + try std.testing.expectEqual(Transport.sse, try selectedTransport()); + } +} + +test "OpenAI Codex SSE fallback requires definitely unsent delivery" { + try std.testing.expect(allowsSseFallback(error.WebSocketUpgradeRejected, .definitely_unsent)); + try std.testing.expect(!allowsSseFallback(error.WebSocketUpgradeRejected, .possibly_sent)); + try std.testing.expect(!allowsSseFallback(error.Cancelled, .definitely_unsent)); + try std.testing.expect(!allowsSseFallback(error.OutOfMemory, .definitely_unsent)); +} + test "OpenAI Codex request uses Responses input and converts AI SDK tool schemas" { const read_file_schema = model_tool_schema.FunctionSchema{ .name = "read_file", @@ -489,14 +743,17 @@ test "OpenAI Codex request uses Responses input and converts AI SDK tool schemas }, .{ .role = .tool, .tool_call_id = "call_1", .tool_name = "read_file", .content = "contents" }, }; - const body = try buildRequest(std.testing.allocator, .{ + const request: stream_provider.RequestData = .{ .model = "gpt-5.4", .messages = &messages, .tools = .{ .additional_functions = &.{read_file_schema} }, .tool_choice = .auto, .provider_options = .{ .reasoning = types.ReasoningEffort.literal("high"), .fast = true }, - }); + }; + const body = try buildRequest(std.testing.allocator, request); defer std.testing.allocator.free(body); + const websocket_body = try buildWebSocketRequest(std.testing.allocator, request); + defer std.testing.allocator.free(websocket_body); try std.testing.expect(std.mem.find(u8, body, "\"model\":\"gpt-5.4\"") != null); try std.testing.expect(std.mem.find(u8, body, "\"instructions\":\"Be concise.\"") != null); @@ -506,6 +763,10 @@ test "OpenAI Codex request uses Responses input and converts AI SDK tool schemas try std.testing.expect(std.mem.find(u8, body, "\"reasoning\":{\"effort\":\"high\"") != null); try std.testing.expect(std.mem.find(u8, body, "\"service_tier\":\"priority\"") != null); try std.testing.expect(std.mem.find(u8, body, "\"max_output_tokens\"") == null); + try std.testing.expect(std.mem.find(u8, websocket_body, "\"type\":\"response.create\"") != null); + try std.testing.expect(std.mem.find(u8, websocket_body, "\"encrypted_content\":\"opaque\"") != null); + try std.testing.expect(std.mem.find(u8, websocket_body, "\"parameters\":{\"type\":\"object\",\"properties\":{}}") != null); + try std.testing.expect(std.mem.find(u8, websocket_body, "\"reasoning\":{\"effort\":\"high\"") != null); } fn makeSizedProviderState(alloc: Allocator, size: usize) ![]u8 { @@ -743,6 +1004,120 @@ test "OpenAI Codex SSE maps text reasoning tools and usage" { try std.testing.expectEqual(types.ProviderFinishReason.tool_calls, completion.finish_reason.?); } +test "OpenAI Codex SSE and WebSocket reducers preserve callback order and completion data" { + const raw_events = [_][]const u8{ + "{\"type\":\"response.output_item.added\",\"output_index\":0,\"item\":{\"type\":\"reasoning\"}}", + "{\"type\":\"response.reasoning_summary_text.delta\",\"output_index\":0,\"delta\":\"thinking\"}", + "{\"type\":\"response.output_item.done\",\"output_index\":0,\"item\":{\"id\":\"rs_1\",\"type\":\"reasoning\",\"summary\":[],\"encrypted_content\":\"opaque\"}}", + "{\"type\":\"response.output_text.delta\",\"output_index\":1,\"delta\":\"hello\"}", + "{\"type\":\"response.output_item.added\",\"output_index\":2,\"item\":{\"type\":\"function_call\",\"call_id\":\"call_1\",\"name\":\"read_file\"}}", + "{\"type\":\"response.function_call_arguments.delta\",\"output_index\":2,\"delta\":\"{\\\"path\\\":\\\"README.md\\\"}\"}", + "{\"type\":\"response.completed\",\"response\":{\"id\":\"response_1\",\"status\":\"completed\",\"usage\":{\"input_tokens\":10,\"output_tokens\":4}}}", + }; + const Capture = struct { + events: std.Io.Writer.Allocating = .init(std.testing.allocator), + failed: bool = false, + + fn deinit(self: *@This()) void { + self.events.deinit(); + } + + fn emit(raw: *anyopaque, event: stream_provider.Event) void { + const self: *@This() = @ptrCast(@alignCast(raw)); + const writer = &self.events.writer; + switch (event) { + .content_delta => |value| { + writer.print("content:{s}|", .{value}) catch { + self.failed = true; + }; + }, + .reasoning_delta => |value| { + writer.print("reasoning:{s}|", .{value}) catch { + self.failed = true; + }; + }, + .tool_started => |tool| { + writer.print("tool:{s}:{s}|", .{ tool.id, tool.name }) catch { + self.failed = true; + }; + }, + .tool_input_delta => |value| { + writer.print("input:{s}|", .{value}) catch { + self.failed = true; + }; + }, + } + } + }; + + var sse_body: std.Io.Writer.Allocating = .init(std.testing.allocator); + defer sse_body.deinit(); + for (raw_events) |raw_event| try sse_body.writer.print("data: {s}\n\n", .{raw_event}); + var sse_capture: Capture = .{}; + defer sse_capture.deinit(); + var sse_events = stream_provider.EventSink{ .context = &sse_capture, .emit_fn = Capture.emit }; + var sse_reader: std.Io.Reader = .fixed(sse_body.written()); + var sse_cancelled = std.atomic.Value(bool).init(false); + const sse_completion = try consumeSse( + std.testing.allocator, + &sse_reader, + &sse_events, + EventBridge.content, + EventBridge.toolStart, + EventBridge.reasoning, + EventBridge.toolInput, + &sse_cancelled, + null, + .{}, + ); + defer freeOpenAICodexTestCompletion(sse_completion); + + var websocket_capture: Capture = .{}; + defer websocket_capture.deinit(); + const websocket_events = stream_provider.EventSink{ .context = &websocket_capture, .emit_fn = Capture.emit }; + var websocket_cancelled = std.atomic.Value(bool).init(false); + var websocket_reducer = responses_protocol.Reducer.init(std.testing.allocator); + defer websocket_reducer.deinit(std.testing.allocator); + var mutable_websocket_events = websocket_events; + const websocket_limits = codexStreamLimits(.{}); + for (raw_events) |raw_event| { + if (try websocket_reducer.applyJson( + std.testing.allocator, + raw_event, + .{ + .context = &mutable_websocket_events, + .on_content = EventBridge.content, + .on_tool_start = EventBridge.toolStart, + .on_reasoning = EventBridge.reasoning, + .on_tool_input = EventBridge.toolInput, + }, + &websocket_cancelled, + null, + websocket_limits, + )) break; + } + const websocket_completion = try websocket_reducer.finish( + std.testing.allocator, + &websocket_cancelled, + websocket_limits, + ); + defer freeOpenAICodexTestCompletion(websocket_completion); + + try std.testing.expect(!sse_capture.failed); + try std.testing.expect(!websocket_capture.failed); + try std.testing.expectEqualStrings(sse_capture.events.written(), websocket_capture.events.written()); + try std.testing.expectEqualStrings(sse_completion.content.?, websocket_completion.content.?); + try std.testing.expectEqualStrings(sse_completion.generation_id.?, websocket_completion.generation_id.?); + try std.testing.expectEqualStrings(sse_completion.provider_state_json.?, websocket_completion.provider_state_json.?); + try std.testing.expectEqual(sse_completion.usage.input_tokens, websocket_completion.usage.input_tokens); + try std.testing.expectEqual(sse_completion.usage.output_tokens, websocket_completion.usage.output_tokens); + try std.testing.expectEqual(sse_completion.finish_reason, websocket_completion.finish_reason); + try std.testing.expectEqual(@as(usize, 1), websocket_completion.tool_calls.len); + try std.testing.expectEqualStrings(sse_completion.tool_calls[0].id, websocket_completion.tool_calls[0].id); + try std.testing.expectEqualStrings(sse_completion.tool_calls[0].name, websocket_completion.tool_calls[0].name); + try std.testing.expectEqualStrings(sse_completion.tool_calls[0].arguments_json, websocket_completion.tool_calls[0].arguments_json); +} + fn consumeOpenAICodexTestSse(sse_text: []const u8, limits: CodexLimits) !types.ModelCompletion { var reader: std.Io.Reader = .fixed(sse_text); var cancelled = std.atomic.Value(bool).init(false); diff --git a/src/gateway/openai_codex_websocket.zig b/src/gateway/openai_codex_websocket.zig new file mode 100644 index 000000000..e63a9920e --- /dev/null +++ b/src/gateway/openai_codex_websocket.zig @@ -0,0 +1,268 @@ +//! Codex-specific WebSocket adapter. +//! +//! This module owns connection reuse, continuation, event reduction, and the +//! WebSocket request lifecycle. The base provider supplies only the shared +//! Responses request serializers and connection credentials. + +const std = @import("std"); +const image_attachments = @import("../core/images/image_attachments.zig"); +const stream_provider = @import("../core/agent/stream_provider.zig"); +const types = @import("../core/shared/types.zig"); +const debug_trace = @import("../core/shared/debug_trace.zig"); +const responses_protocol = @import("responses_protocol.zig"); +const websocket_transport = @import("websocket_transport.zig"); +const codex_websocket_session = @import("codex_websocket_session.zig"); + +const Allocator = std.mem.Allocator; + +pub const ConnectionConfig = struct { + endpoint: []const u8, + authorization: []const u8, + account_id: []const u8, +}; + +pub const Serializer = struct { + build_input: *const fn ( + alloc: Allocator, + messages: []const types.ChatMessage, + images: ?[]const image_attachments.VerifiedSnapshot, + ) anyerror![]u8, + build_request: *const fn ( + alloc: Allocator, + request: stream_provider.RequestData, + input: []const u8, + previous_response_id: ?[]const u8, + ) anyerror![]u8, + build_durable_input: *const fn ( + alloc: Allocator, + messages: []const types.ChatMessage, + ) anyerror![]u8, +}; + +pub fn shutdown() void { + codex_websocket_session.shutdown(); +} + +pub fn stream( + alloc: Allocator, + request: stream_provider.ModelRequest, + config: ConnectionConfig, + serializer: Serializer, + stream_limits: responses_protocol.StreamLimits, +) !stream_provider.Result { + if (request.cancel_flag.load(.seq_cst)) return error.Cancelled; + + const full_input = try serializer.build_input(alloc, request.messages, request.verified_images); + defer alloc.free(full_input); + const shape_payload = try serializer.build_request(alloc, request.data(), "", null); + defer alloc.free(shape_payload); + var shape: [std.crypto.hash.sha2.Sha256.digest_length]u8 = undefined; + std.crypto.hash.sha2.Sha256.hash(shape_payload, &shape, .{}); + + var reducer = responses_protocol.Reducer.init(alloc); + defer reducer.deinit(alloc); + var bridge = WebSocketBridge{ + .alloc = alloc, + .reducer = &reducer, + .events = request.events, + .cancel_flag = request.cancel_flag, + .content_capture_limit = request.content_capture_limit, + .stream_limits = stream_limits, + }; + try request.admission.admit(); + + var continuation_recovery_attempted = false; + while (true) { + const checkout = try codex_websocket_session.acquire(alloc, .{ + .session_id = request.session_id, + .account_id = config.account_id, + .model = request.model, + .endpoint = config.endpoint, + .authorization = config.authorization, + .deadline = request.deadline, + .cancel_flag = request.cancel_flag, + .delivery = request.delivery, + .continuation_input = if (continuation_recovery_attempted) null else full_input, + .continuation_shape = if (continuation_recovery_attempted) null else shape, + }); + debug_trace.eventf("codex.ws", "turn", request.trace_ctx, "lane={d} retained={d} reused={d} handshake_ms={d} health={d} auth=chatgpt_subscription", .{ + checkout.slot orelse std.math.maxInt(usize), + @as(u8, @intFromBool(checkout.retained)), + @as(u8, @intFromBool(checkout.reused)), + checkout.handshake_ms, + checkout.health_failures, + }); + const continued = if (continuation_recovery_attempted) + null + else + codex_websocket_session.continuation(checkout.slot, full_input, shape); + const payload = try serializer.build_request( + alloc, + request.data(), + if (continued) |value| value.delta_input else full_input, + if (continued) |value| value.previous_response_id else null, + ); + defer alloc.free(payload); + debug_trace.eventf("codex.ws", "continuation", request.trace_ctx, "used={d} delta_bytes={d} recovery={d}", .{ + @as(u8, @intFromBool(continued != null)), + if (continued) |value| value.delta_input.len else full_input.len, + @as(u8, @intFromBool(continuation_recovery_attempted)), + }); + websocket_transport.streamOn(checkout.connection, alloc, .{ + .endpoint = config.endpoint, + .authorization = config.authorization, + .account_id = config.account_id, + .session_id = request.session_id, + .payload = payload, + .deadline = request.deadline, + .cancel_flag = request.cancel_flag, + .delivery = request.delivery, + }, &bridge, WebSocketBridge.event) catch |err| { + codex_websocket_session.release(checkout, .failed); + if (err == error.PreviousResponseNotFound and continued != null and !continuation_recovery_attempted) { + continuation_recovery_attempted = true; + reducer.deinit(alloc); + reducer = responses_protocol.Reducer.init(alloc); + continue; + } + debug_trace.eventf("codex.ws", "poison", request.trace_ctx, "reason={s} close={d}", .{ failureReason(err), @as(u16, 0) }); + return err; + }; + const completion = reducer.finish(alloc, request.cancel_flag, bridge.stream_limits) catch |err| { + codex_websocket_session.release(checkout, .failed); + debug_trace.eventf("codex.ws", "poison", request.trace_ctx, "reason=protocol close={d}", .{@as(u16, 0)}); + return mapReducerError(err); + }; + recordContinuation(alloc, checkout.slot, request, completion, full_input, shape, serializer); + codex_websocket_session.release(checkout, .completed); + return .{ .completed = .{ + .completion = completion, + .usage = .{ .unavailable = .possibly_billed }, + .ownership = .owned, + } }; + } +} + +fn recordContinuation( + alloc: Allocator, + slot: ?usize, + request: stream_provider.ModelRequest, + completion: types.ModelCompletion, + full_input: []const u8, + shape: [std.crypto.hash.sha2.Sha256.digest_length]u8, + serializer: Serializer, +) void { + const response_id = completion.generation_id orelse return; + const response_message = [_]types.ChatMessage{.{ + .role = .assistant, + .content = completion.content, + .tool_calls = completion.tool_calls, + .provider_state_json = completion.provider_state_json, + }}; + const response_input = serializer.build_input(alloc, &response_message, null) catch return; + defer alloc.free(response_input); + const baseline = buildContinuationBaseline(alloc, full_input, response_input) catch return; + defer alloc.free(baseline); + const durable_full_input = serializer.build_durable_input(alloc, request.messages) catch return; + defer alloc.free(durable_full_input); + const durable_response_input = serializer.build_durable_input(alloc, &response_message) catch return; + defer alloc.free(durable_response_input); + const durable_baseline = buildContinuationBaseline(alloc, durable_full_input, durable_response_input) catch return; + defer alloc.free(durable_baseline); + codex_websocket_session.recordCompletion( + slot, + response_id, + baseline, + durable_baseline, + shape, + ); +} + +fn buildContinuationBaseline( + alloc: Allocator, + full_input: []const u8, + response_input: []const u8, +) ![]u8 { + var out: std.Io.Writer.Allocating = .init(alloc); + errdefer out.deinit(); + try out.writer.writeAll(full_input); + if (response_input.len > 0) { + if (out.written().len > 0) try out.writer.writeByte(','); + try out.writer.writeAll(response_input); + } + return out.toOwnedSlice(); +} + +fn failureReason(err: anyerror) []const u8 { + return switch (err) { + error.Cancelled => "cancel", + error.Timeout => "timeout", + error.WebSocketPolicyClosed => "policy", + error.WebSocketUnexpectedBinary => "binary", + error.WebSocketProtocolViolation, error.WebSocketInvalidUtf8 => "protocol", + error.WebSocketUpgradeRejected, error.WebSocketAcceptInvalid => "auth", + else => "close", + }; +} + +const WebSocketBridge = struct { + alloc: Allocator, + reducer: *responses_protocol.Reducer, + events: stream_provider.EventSink, + cancel_flag: *std.atomic.Value(bool), + content_capture_limit: ?usize, + stream_limits: responses_protocol.StreamLimits, + + fn event(raw: *anyopaque, json_text: []const u8) !bool { + const self: *@This() = @ptrCast(@alignCast(raw)); + return self.reducer.applyJson( + self.alloc, + json_text, + .{ + .context = &self.events, + .on_content = EventBridge.content, + .on_tool_start = EventBridge.toolStart, + .on_reasoning = EventBridge.reasoning, + .on_tool_input = EventBridge.toolInput, + }, + self.cancel_flag, + self.content_capture_limit, + self.stream_limits, + ) catch |err| return mapReducerError(err); + } +}; + +const EventBridge = struct { + fn sink(raw: *anyopaque) *stream_provider.EventSink { + return @ptrCast(@alignCast(raw)); + } + + fn content(raw: *anyopaque, chunk: []const u8) void { + sink(raw).emit(.{ .content_delta = chunk }); + } + + fn reasoning(raw: *anyopaque, chunk: []const u8) void { + sink(raw).emit(.{ .reasoning_delta = chunk }); + } + + fn toolInput(raw: *anyopaque, chunk: []const u8) void { + sink(raw).emit(.{ .tool_input_delta = chunk }); + } + + fn toolStart(raw: *anyopaque, id: []const u8, name: []const u8, label: ?[]const u8) void { + sink(raw).emit(.{ .tool_started = .{ .id = id, .name = name, .label = label } }); + } +}; + +fn mapReducerError(err: anyerror) anyerror { + return switch (err) { + error.InvalidEvent => error.InvalidOpenAICodexSseEvent, + error.PreviousResponseNotFound => error.PreviousResponseNotFound, + error.ResponseFailed => error.OpenAICodexResponseFailed, + error.StreamIncomplete => error.OpenAICodexStreamIncomplete, + error.ToolCallLimitExceeded => error.OpenAICodexToolCallLimitExceeded, + error.ToolArgumentsTooLarge => error.OpenAICodexToolArgumentsTooLarge, + error.ResourceLimitExceeded => error.OpenAICodexResourceLimitExceeded, + else => err, + }; +} diff --git a/src/gateway/responses_protocol.zig b/src/gateway/responses_protocol.zig index 4707a13d7..186371d5a 100644 --- a/src/gateway/responses_protocol.zig +++ b/src/gateway/responses_protocol.zig @@ -378,6 +378,21 @@ pub const Reducer = struct { } else if (std.mem.eql(u8, event_type, "response.failed") or std.mem.eql(u8, event_type, "error")) { + const error_value = parsed.value.object.get("error"); + const response_value = parsed.value.object.get("response"); + const response_error = if (response_value != null and response_value.? == .object) + response_value.?.object.get("error") + else + null; + const code = if (error_value != null and error_value.? == .object) + stringField(error_value.?.object, "code") + else if (response_error != null and response_error.? == .object) + stringField(response_error.?.object, "code") + else + stringField(parsed.value.object, "code"); + if (code) |value| if (std.mem.eql(u8, value, "previous_response_not_found")) { + return error.PreviousResponseNotFound; + }; return error.ResponseFailed; } return false; @@ -738,6 +753,39 @@ test "Responses usage projection retains optional cached and reasoning detail" { try std.testing.expectEqual(@as(?u64, 3), usage.reasoning_tokens); } +test "Responses protocol distinguishes missing WebSocket continuation state" { + const Capture = struct { + fn content(_: *anyopaque, _: []const u8) void {} + fn toolStart(_: *anyopaque, _: []const u8, _: []const u8, _: ?[]const u8) void {} + }; + var reducer = Reducer.init(std.testing.allocator); + defer reducer.deinit(std.testing.allocator); + var cancelled = std.atomic.Value(bool).init(false); + var context: u8 = 0; + try std.testing.expectError( + error.PreviousResponseNotFound, + reducer.applyJson( + std.testing.allocator, + "{\"type\":\"error\",\"error\":{\"code\":\"previous_response_not_found\",\"message\":\"expired\"}}", + .{ + .context = &context, + .on_content = Capture.content, + .on_tool_start = Capture.toolStart, + }, + &cancelled, + null, + .{ + .aggregate_bytes = 4096, + .events = 8, + .tool_calls = 8, + .tool_identity_bytes = 1024, + .tool_arguments_bytes = 4096, + .provider_state_bytes = 4096, + }, + ), + ); +} + test "Responses protocol owns one subscription billing projection" { const alloc = std.testing.allocator; const billing = (try buildSubscriptionBilling( diff --git a/src/gateway/websocket_transport.zig b/src/gateway/websocket_transport.zig new file mode 100644 index 000000000..96b3e8d5b --- /dev/null +++ b/src/gateway/websocket_transport.zig @@ -0,0 +1,1145 @@ +const std = @import("std"); +const io_mod = @import("../core/shared/io.zig"); +const gateway_client = @import("client.zig"); + +const Allocator = std.mem.Allocator; + +pub const max_frame_bytes: usize = 4 * 1024 * 1024; +pub const max_message_bytes: usize = 64 * 1024 * 1024; + +pub const Error = error{ + WebSocketUpgradeRejected, + WebSocketAcceptInvalid, + WebSocketProtocolViolation, + WebSocketUnexpectedBinary, + WebSocketMessageTooLarge, + WebSocketInvalidUtf8, + WebSocketPolicyClosed, + WebSocketClosedBeforeCompletion, +}; + +pub const EventHandler = *const fn (context: *anyopaque, json: []const u8) anyerror!bool; + +const default_connect_timeout_ms: i64 = 30_000; +const default_event_idle_timeout_ms: i64 = 30_000; +const connect_timeout_env = "FX_CODEX_WEBSOCKET_CONNECT_TIMEOUT_MS"; +const event_idle_timeout_env = "FX_CODEX_WEBSOCKET_EVENT_IDLE_TIMEOUT_MS"; + +pub const ConnectArgs = struct { + endpoint: []const u8, + authorization: []const u8, + account_id: []const u8, + session_id: ?[]const u8, + deadline: ?std.Io.Clock.Timestamp, + cancel_flag: *std.atomic.Value(bool), + delivery: *gateway_client.DeliveryCertainty, +}; + +pub const StreamArgs = struct { + endpoint: []const u8, + authorization: []const u8, + account_id: []const u8, + session_id: ?[]const u8, + payload: []const u8, + deadline: ?std.Io.Clock.Timestamp, + cancel_flag: *std.atomic.Value(bool), + delivery: *gateway_client.DeliveryCertainty, +}; + +pub const Request = StreamArgs; + +pub const Connection = struct { + alloc: Allocator, + client: std.http.Client, + request: std.http.Client.Request, + opened_at_ms: i64, + close_sent: bool = false, + watcher_done: std.atomic.Value(bool) = std.atomic.Value(bool).init(true), + timeout_fired: std.atomic.Value(bool) = std.atomic.Value(bool).init(false), + last_progress_ms: std.atomic.Value(i64), + watcher: ?std.Thread = null, + + fn socket(self: *Connection) !*std.http.Client.Connection { + return self.request.connection orelse error.WebSocketConnectionMissing; + } + + fn startWatcher(self: *Connection, cancel_flag: *std.atomic.Value(bool), deadline: ?std.Io.Clock.Timestamp) !void { + self.watcher_done.store(false, .seq_cst); + self.timeout_fired.store(false, .seq_cst); + self.last_progress_ms.store(io_mod.milliTimestamp(), .seq_cst); + const http_connection = try self.socket(); + self.watcher = try spawnConnectionWatcher( + &self.watcher_done, + cancel_flag, + deadline, + &self.timeout_fired, + &self.last_progress_ms, + try eventIdleTimeoutMs(), + http_connection.stream_writer.stream, + ); + } + + fn stopWatcher(self: *Connection) void { + self.watcher_done.store(true, .seq_cst); + if (self.watcher) |thread| thread.join(); + self.watcher = null; + } +}; + +const OpenedRequest = struct { + request: ?std.http.Client.Request, + + pub fn deinit(self: *OpenedRequest, _: Allocator) void { + if (self.request) |*request| request.deinit(); + self.request = null; + } + + fn take(self: *OpenedRequest) std.http.Client.Request { + const request = self.request.?; + self.request = null; + return request; + } +}; + +const OpenWebSocketOperation = struct { + client: *std.http.Client, + uri: std.Uri, + authorization: []const u8, + headers: []const std.http.Header, + + pub fn run(self: *@This()) !OpenedRequest { + return .{ .request = try self.client.request(.GET, self.uri, .{ + .headers = .{ + .authorization = .{ .override = self.authorization }, + .connection = .{ .override = "Upgrade" }, + .accept_encoding = .omit, + }, + .extra_headers = self.headers, + .keep_alive = false, + .redirect_behavior = .unhandled, + }) }; + } +}; + +fn positiveTimeoutFromEnv(name: []const u8, fallback: i64) !i64 { + const value = io_mod.getenv(name) orelse return fallback; + const parsed = std.fmt.parseInt(i64, value, 10) catch return error.InvalidOpenAICodexTransport; + if (parsed <= 0) return error.InvalidOpenAICodexTransport; + return parsed; +} + +fn connectTimeoutMs() !i64 { + return positiveTimeoutFromEnv(connect_timeout_env, default_connect_timeout_ms); +} + +fn eventIdleTimeoutMs() !i64 { + return positiveTimeoutFromEnv(event_idle_timeout_env, default_event_idle_timeout_ms); +} + +pub fn connect(alloc: Allocator, args: ConnectArgs) !*Connection { + if (args.cancel_flag.load(.seq_cst)) return error.Cancelled; + const uri = try std.Uri.parse(args.endpoint); + var nonce: [16]u8 = undefined; + try io_mod.getIo().randomSecure(&nonce); + var key_buffer: [std.base64.standard.Encoder.calcSize(nonce.len)]u8 = undefined; + _ = std.base64.standard.Encoder.encode(&key_buffer, &nonce); + var accept_buffer: [std.base64.standard.Encoder.calcSize(std.crypto.hash.Sha1.digest_length)]u8 = undefined; + const expected_accept = websocketAccept(&key_buffer, &accept_buffer); + + var extra_headers: [7]std.http.Header = undefined; + var count: usize = 0; + extra_headers[count] = .{ .name = "chatgpt-account-id", .value = args.account_id }; + count += 1; + extra_headers[count] = .{ .name = "originator", .value = "fx" }; + count += 1; + extra_headers[count] = .{ .name = "OpenAI-Beta", .value = "responses_websockets=2026-02-06" }; + count += 1; + extra_headers[count] = .{ .name = "Upgrade", .value = "websocket" }; + count += 1; + extra_headers[count] = .{ .name = "Sec-WebSocket-Version", .value = "13" }; + count += 1; + extra_headers[count] = .{ .name = "Sec-WebSocket-Key", .value = &key_buffer }; + count += 1; + if (args.session_id) |session_id| if (session_id.len > 0) { + extra_headers[count] = .{ .name = "session-id", .value = session_id }; + count += 1; + }; + + const connection = try alloc.create(Connection); + errdefer alloc.destroy(connection); + connection.* = undefined; + connection.alloc = alloc; + connection.client = .{ .allocator = alloc, .io = io_mod.getIo() }; + errdefer connection.client.deinit(); + var open_operation = OpenWebSocketOperation{ + .client = &connection.client, + .uri = uri, + .authorization = args.authorization, + .headers = extra_headers[0..count], + }; + var connect_deadline = std.Io.Clock.Timestamp.fromNow(io_mod.getIo(), .{ + .clock = .awake, + .raw = .fromMilliseconds(try connectTimeoutMs()), + }); + if (args.deadline) |deadline| { + if (std.Io.Clock.Timestamp.compare(deadline, .lt, connect_deadline)) connect_deadline = deadline; + } + var opened = try gateway_client.runBoundedHttpOperation( + OpenedRequest, + alloc, + args.cancel_flag, + connect_deadline, + &open_operation, + ); + errdefer opened.deinit(alloc); + connection.request = opened.take(); + errdefer connection.request.deinit(); + connection.opened_at_ms = io_mod.milliTimestamp(); + connection.close_sent = false; + connection.watcher_done = std.atomic.Value(bool).init(true); + connection.timeout_fired = std.atomic.Value(bool).init(false); + connection.last_progress_ms = std.atomic.Value(i64).init(connection.opened_at_ms); + connection.watcher = null; + + connection.request.sendBodiless() catch |err| { + if (args.cancel_flag.load(.seq_cst)) return error.Cancelled; + return err; + }; + const response = connection.request.receiveHead(&.{}) catch |err| { + if (args.cancel_flag.load(.seq_cst)) return error.Cancelled; + return err; + }; + if (response.head.status != .switching_protocols) return error.WebSocketUpgradeRejected; + if (!hasHeader(response.head, "upgrade", "websocket") or + !hasTokenHeader(response.head, "connection", "upgrade") or + !hasHeader(response.head, "sec-websocket-accept", expected_accept)) + { + return error.WebSocketAcceptInvalid; + } + _ = try connection.socket(); + return connection; +} + +fn operationError(connection: *Connection, cancel_flag: *std.atomic.Value(bool), err: anyerror) anyerror { + if (cancel_flag.load(.seq_cst)) return error.Cancelled; + if (connection.timeout_fired.load(.seq_cst)) return error.Timeout; + return err; +} + +pub fn ping( + connection: *Connection, + cancel_flag: *std.atomic.Value(bool), + deadline: ?std.Io.Clock.Timestamp, + _: *gateway_client.DeliveryCertainty, +) !void { + if (cancel_flag.load(.seq_cst)) return error.Cancelled; + try connection.startWatcher(cancel_flag, deadline); + defer connection.stopWatcher(); + const socket = try connection.socket(); + const writer = socket.writer(); + writeFrame(writer, .ping, &.{}) catch |err| return operationError(connection, cancel_flag, err); + socket.flush() catch |err| return operationError(connection, cancel_flag, err); + while (true) { + const frame = readFrame(connection.alloc, connection.request.reader.in) catch |err| return operationError(connection, cancel_flag, err); + defer connection.alloc.free(frame.payload); + connection.last_progress_ms.store(io_mod.milliTimestamp(), .seq_cst); + switch (frame.opcode) { + .pong => return, + .ping => { + writeFrame(writer, .pong, frame.payload) catch |err| return operationError(connection, cancel_flag, err); + socket.flush() catch |err| return operationError(connection, cancel_flag, err); + }, + .close => return closeError(try validateClosePayload(frame.payload)), + .binary => return error.WebSocketUnexpectedBinary, + else => return error.WebSocketProtocolViolation, + } + } +} + +pub fn streamOn( + connection: *Connection, + alloc: Allocator, + request: StreamArgs, + context: *anyopaque, + on_event: EventHandler, +) !void { + if (request.cancel_flag.load(.seq_cst)) return error.Cancelled; + try connection.startWatcher(request.cancel_flag, request.deadline); + defer connection.stopWatcher(); + const socket = try connection.socket(); + const reader = connection.request.reader.in; + const writer = socket.writer(); + var succeeded = false; + defer if (!succeeded) closeAfterFailure(connection, request.cancel_flag); + request.delivery.markPossiblySent(); + writeFrame(writer, .text, request.payload) catch |err| return operationError(connection, request.cancel_flag, err); + socket.flush() catch |err| return operationError(connection, request.cancel_flag, err); + connection.last_progress_ms.store(io_mod.milliTimestamp(), .seq_cst); + + var message: std.ArrayList(u8) = .empty; + defer message.deinit(alloc); + var fragmented_opcode: ?Opcode = null; + while (true) { + if (request.cancel_flag.load(.seq_cst)) return error.Cancelled; + const frame = readFrame(alloc, reader) catch |err| return operationError(connection, request.cancel_flag, err); + connection.last_progress_ms.store(io_mod.milliTimestamp(), .seq_cst); + defer alloc.free(frame.payload); + switch (frame.opcode) { + .ping => { + writeFrame(writer, .pong, frame.payload) catch |err| return operationError(connection, request.cancel_flag, err); + socket.flush() catch |err| return operationError(connection, request.cancel_flag, err); + }, + .pong => {}, + .close => return closeError(try validateClosePayload(frame.payload)), + .binary => return error.WebSocketUnexpectedBinary, + .continuation => { + if (fragmented_opcode == null) return error.WebSocketProtocolViolation; + try appendMessage(&message, alloc, frame.payload); + if (!frame.fin) continue; + const opcode = fragmented_opcode.?; + fragmented_opcode = null; + if (opcode != .text) return error.WebSocketUnexpectedBinary; + if (try dispatchTextMessage(context, on_event, message.items)) { + succeeded = true; + return; + } + message.clearRetainingCapacity(); + }, + .text => { + if (fragmented_opcode != null) return error.WebSocketProtocolViolation; + try appendMessage(&message, alloc, frame.payload); + if (!frame.fin) { + fragmented_opcode = .text; + continue; + } + if (try dispatchTextMessage(context, on_event, message.items)) { + succeeded = true; + return; + } + message.clearRetainingCapacity(); + }, + } + } +} + +fn closeAfterFailure(connection: *Connection, cancel_flag: *std.atomic.Value(bool)) void { + if (connection.close_sent or cancel_flag.load(.seq_cst)) return; + const socket = connection.socket() catch return; + writeFrame(socket.writer(), .close, &.{ 0x03, 0xe8 }) catch return; + socket.flush() catch {}; + connection.close_sent = true; +} + +fn closeChecked( + connection: *Connection, + alloc: Allocator, + cancel_flag: *std.atomic.Value(bool), + deadline: ?std.Io.Clock.Timestamp, +) !void { + connection.stopWatcher(); + defer { + if (connection.request.connection) |http_connection| http_connection.closing = true; + connection.request.deinit(); + connection.client.deinit(); + alloc.destroy(connection); + } + if (!connection.close_sent) { + try connection.startWatcher(cancel_flag, deadline); + defer connection.stopWatcher(); + const http_connection = try connection.socket(); + try closeAfterCompletion( + alloc, + connection.request.reader.in, + http_connection.writer(), + http_connection, + cancel_flag, + &connection.timeout_fired, + ); + connection.close_sent = true; + } +} + +pub fn close(connection: *Connection, alloc: Allocator) void { + var cancelled = std.atomic.Value(bool).init(false); + const deadline = std.Io.Clock.Timestamp.fromNow(io_mod.getIo(), .{ + .clock = .awake, + .raw = .fromMilliseconds(1_000), + }); + closeChecked(connection, alloc, &cancelled, deadline) catch {}; +} + +/// Opens one socket, sends one request, and consumes one terminal response. +pub fn stream( + alloc: Allocator, + request: Request, + context: *anyopaque, + on_event: EventHandler, +) !void { + const connection = try connect(alloc, .{ + .endpoint = request.endpoint, + .authorization = request.authorization, + .account_id = request.account_id, + .session_id = request.session_id, + .deadline = request.deadline, + .cancel_flag = request.cancel_flag, + .delivery = request.delivery, + }); + var owned = true; + defer if (owned) close(connection, alloc); + try streamOn(connection, alloc, request, context, on_event); + owned = false; + return closeChecked(connection, alloc, request.cancel_flag, request.deadline); +} + +const Opcode = enum(u4) { continuation = 0, text = 1, binary = 2, close = 8, ping = 9, pong = 10 }; +const Frame = struct { fin: bool, opcode: Opcode, payload: []u8 }; + +fn websocketAccept(key: []const u8, output: []u8) []const u8 { + var hash = std.crypto.hash.Sha1.init(.{}); + hash.update(key); + hash.update("258EAFA5-E914-47DA-95CA-C5AB0DC85B11"); + var digest: [std.crypto.hash.Sha1.digest_length]u8 = undefined; + hash.final(&digest); + _ = std.base64.standard.Encoder.encode(output, &digest); + return output; +} + +fn hasHeader(head: std.http.Client.Response.Head, name: []const u8, expected: []const u8) bool { + var it = head.iterateHeaders(); + while (it.next()) |header| { + if (std.ascii.eqlIgnoreCase(header.name, name) and std.mem.eql(u8, header.value, expected)) return true; + } + return false; +} + +fn hasTokenHeader(head: std.http.Client.Response.Head, name: []const u8, token: []const u8) bool { + var it = head.iterateHeaders(); + while (it.next()) |header| { + if (!std.ascii.eqlIgnoreCase(header.name, name)) continue; + var tokens = std.mem.splitScalar(u8, header.value, ','); + while (tokens.next()) |candidate| if (std.ascii.eqlIgnoreCase(std.mem.trim(u8, candidate, " \t"), token)) return true; + } + return false; +} + +fn appendMessage(message: *std.ArrayList(u8), alloc: Allocator, payload: []const u8) !void { + if (payload.len > max_message_bytes -| message.items.len) return error.WebSocketMessageTooLarge; + try message.appendSlice(alloc, payload); +} + +fn dispatchTextMessage(context: *anyopaque, on_event: EventHandler, message: []const u8) !bool { + if (!std.unicode.utf8ValidateSlice(message)) return error.WebSocketInvalidUtf8; + return on_event(context, message); +} + +fn closeAfterCompletion( + alloc: Allocator, + reader: *std.Io.Reader, + writer: *std.Io.Writer, + connection: anytype, + cancel_flag: *std.atomic.Value(bool), + timeout_fired: *std.atomic.Value(bool), +) !void { + try writeFrame(writer, .close, &.{ 0x03, 0xe8 }); + try connection.flush(); + const frame = readFrame(alloc, reader) catch |err| { + if (cancel_flag.load(.seq_cst)) return error.Cancelled; + if (timeout_fired.load(.seq_cst)) return error.Timeout; + return err; + }; + defer alloc.free(frame.payload); + switch (frame.opcode) { + .close => _ = try validateClosePayload(frame.payload), + else => return error.WebSocketProtocolViolation, + } +} + +fn closeError(code: ?u16) anyerror { + if (code == 1008) return error.WebSocketPolicyClosed; + return error.WebSocketClosedBeforeCompletion; +} + +fn validateClosePayload(payload: []const u8) !?u16 { + if (payload.len == 1) return error.WebSocketProtocolViolation; + if (payload.len < 2) return null; + const code = std.mem.readInt(u16, payload[0..2], .big); + if (code < 1000 or code >= 5000 or code == 1004 or code == 1005 or code == 1006 or code == 1015) { + return error.WebSocketProtocolViolation; + } + if (!std.unicode.utf8ValidateSlice(payload[2..])) return error.WebSocketInvalidUtf8; + return code; +} + +const ConnectionWatcher = struct { + fn run( + done: *std.atomic.Value(bool), + cancel_flag: *std.atomic.Value(bool), + deadline: ?std.Io.Clock.Timestamp, + timeout_fired: *std.atomic.Value(bool), + last_progress_ms: *std.atomic.Value(i64), + event_idle_timeout_ms: i64, + socket: std.Io.net.Stream, + ) void { + while (!done.load(.seq_cst)) { + if (cancel_flag.load(.seq_cst)) { + socket.shutdown(io_mod.getIo(), .both) catch {}; + return; + } + if (deadline) |limit| { + const now = std.Io.Clock.Timestamp.now(io_mod.getIo(), .awake); + if (!std.Io.Clock.Timestamp.compare(now, .lt, limit)) { + timeout_fired.store(true, .seq_cst); + socket.shutdown(io_mod.getIo(), .both) catch {}; + return; + } + } + const elapsed_ms = io_mod.milliTimestamp() - last_progress_ms.load(.seq_cst); + if (elapsed_ms >= event_idle_timeout_ms) { + timeout_fired.store(true, .seq_cst); + socket.shutdown(io_mod.getIo(), .both) catch {}; + return; + } + io_mod.sleep(10 * std.time.ns_per_ms); + } + } +}; + +fn spawnConnectionWatcher( + done: *std.atomic.Value(bool), + cancel_flag: *std.atomic.Value(bool), + deadline: ?std.Io.Clock.Timestamp, + timeout_fired: *std.atomic.Value(bool), + last_progress_ms: *std.atomic.Value(i64), + event_idle_timeout_ms: i64, + socket: std.Io.net.Stream, +) !std.Thread { + return std.Thread.spawn(.{}, ConnectionWatcher.run, .{ + done, + cancel_flag, + deadline, + timeout_fired, + last_progress_ms, + event_idle_timeout_ms, + socket, + }); +} + +fn writeFrame(writer: *std.Io.Writer, opcode: Opcode, payload: []const u8) !void { + if (payload.len > max_message_bytes) return error.WebSocketMessageTooLarge; + var mask: [4]u8 = undefined; + try io_mod.getIo().randomSecure(&mask); + try writer.writeByte(0x80 | @as(u8, @intFromEnum(opcode))); + if (payload.len < 126) { + try writer.writeByte(0x80 | @as(u8, @intCast(payload.len))); + } else if (payload.len <= std.math.maxInt(u16)) { + try writer.writeByte(0x80 | 126); + try writer.writeInt(u16, @intCast(payload.len), .big); + } else { + try writer.writeByte(0x80 | 127); + try writer.writeInt(u64, @intCast(payload.len), .big); + } + try writer.writeAll(&mask); + var chunk: [4096]u8 = undefined; + var offset: usize = 0; + while (offset < payload.len) { + const length = @min(chunk.len, payload.len - offset); + for (payload[offset..][0..length], 0..) |byte, index| chunk[index] = byte ^ mask[(offset + index) % mask.len]; + try writer.writeAll(chunk[0..length]); + offset += length; + } +} + +fn readFrame(alloc: Allocator, reader: *std.Io.Reader) !Frame { + const first = try reader.takeByte(); + const second = try reader.takeByte(); + if (second & 0x80 != 0 or first & 0x70 != 0) return error.WebSocketProtocolViolation; + const fin = first & 0x80 != 0; + const opcode = std.enums.fromInt(Opcode, first & 0x0f) orelse return error.WebSocketProtocolViolation; + var length: u64 = second & 0x7f; + if (length == 126) length = try reader.takeInt(u16, .big) else if (length == 127) { + length = try reader.takeInt(u64, .big); + if (length & (@as(u64, 1) << 63) != 0) return error.WebSocketProtocolViolation; + } + if (length > max_frame_bytes) return error.WebSocketMessageTooLarge; + if (@intFromEnum(opcode) >= @intFromEnum(Opcode.close) and (!fin or length > 125)) return error.WebSocketProtocolViolation; + const payload = try alloc.alloc(u8, @intCast(length)); + errdefer alloc.free(payload); + try reader.readSliceAll(payload); + return .{ .fin = fin, .opcode = opcode, .payload = payload }; +} + +test "WebSocket accept matches RFC 6455" { + var output: [28]u8 = undefined; + try std.testing.expectEqualStrings("s3pPLMBiTxaQ9kYGzzhZRbK+xOo=", websocketAccept("dGhlIHNhbXBsZSBub25jZQ==", &output)); +} + +test "fragment aggregation limits message size" { + var message: std.ArrayList(u8) = .empty; + defer message.deinit(std.testing.allocator); + try appendMessage(&message, std.testing.allocator, "hello"); + try std.testing.expectEqualStrings("hello", message.items); +} + +test "extended 127-byte frame keeps the following frame aligned" { + var encoded: std.Io.Writer.Allocating = .init(std.testing.allocator); + defer encoded.deinit(); + try encoded.writer.writeAll(&.{ 0x81, 126, 0, 127 }); + try encoded.writer.splatByteAll('a', 127); + try encoded.writer.writeAll(&.{ 0x81, 2, 'o', 'k' }); + + var reader = std.Io.Reader.fixed(encoded.written()); + const first = try readFrame(std.testing.allocator, &reader); + defer std.testing.allocator.free(first.payload); + try std.testing.expectEqual(Opcode.text, first.opcode); + try std.testing.expectEqual(@as(usize, 127), first.payload.len); + + const second = try readFrame(std.testing.allocator, &reader); + defer std.testing.allocator.free(second.payload); + try std.testing.expectEqual(Opcode.text, second.opcode); + try std.testing.expectEqualStrings("ok", second.payload); +} + +test "text messages reject malformed UTF-8 before event dispatch" { + const Handler = struct { + fn handle(_: *anyopaque, _: []const u8) !bool { + return false; + } + }; + var context: u8 = 0; + try std.testing.expectError( + error.WebSocketInvalidUtf8, + dispatchTextMessage(@ptrCast(&context), Handler.handle, &.{ 0xc3, 0x28 }), + ); +} + +test "close payload rejects reserved codes and malformed UTF-8 reasons" { + try std.testing.expectError(error.WebSocketProtocolViolation, validateClosePayload(&.{ 0x03, 0xed })); + try std.testing.expectError(error.WebSocketInvalidUtf8, validateClosePayload(&.{ 0x03, 0xe8, 0xc3, 0x28 })); + try std.testing.expectEqual(@as(?u16, 1000), try validateClosePayload(&.{ 0x03, 0xe8, 'o', 'k' })); + try std.testing.expectEqual(error.WebSocketPolicyClosed, closeError(1008)); + try std.testing.expectEqual(error.WebSocketClosedBeforeCompletion, closeError(1000)); +} + +const StalledWriteFixture = struct { + io_backend: std.Io.Threaded = .init_single_threaded, + server: std.Io.net.Server, + thread: ?std.Thread = null, + stopping: std.atomic.Value(bool) = .init(false), + upgraded: std.atomic.Value(bool) = .init(false), + failure: ?anyerror = null, + + fn init() !@This() { + var fixture: @This() = .{ .server = undefined }; + var address = try std.Io.net.IpAddress.parse("127.0.0.1", 0); + fixture.server = try address.listen(fixture.io(), .{ .reuse_address = true }); + return fixture; + } + + fn io(self: *@This()) std.Io { + return self.io_backend.io(); + } + + fn endpoint(self: *@This(), buffer: []u8) ![]const u8 { + return std.fmt.bufPrint(buffer, "http://127.0.0.1:{d}/responses", .{self.server.socket.address.getPort()}); + } + + fn start(self: *@This()) !void { + self.thread = try std.Thread.spawn(.{}, run, .{self}); + } + + fn deinit(self: *@This()) void { + self.stopping.store(true, .seq_cst); + if (self.thread) |thread| { + const listener = std.Io.net.Stream{ .socket = self.server.socket }; + listener.shutdown(self.io(), .both) catch {}; + thread.join(); + self.thread = null; + } + self.server.deinit(self.io()); + } + + fn run(self: *@This()) void { + self.runFallible() catch |err| { + if (!self.stopping.load(.seq_cst)) self.failure = err; + }; + } + + fn runFallible(self: *@This()) !void { + const zio = self.io(); + var client_stream = try self.server.accept(zio); + defer client_stream.close(zio); + if (self.stopping.load(.seq_cst)) return; + const receive_buffer: c_int = 1024; + std.posix.setsockopt(client_stream.socket.handle, std.posix.SOL.SOCKET, std.posix.SO.RCVBUF, std.mem.asBytes(&receive_buffer)) catch {}; + + var socket_buffer: [4096]u8 = undefined; + var reader = client_stream.reader(zio, &socket_buffer); + var request: [16 * 1024]u8 = undefined; + var request_len: usize = 0; + while (request_len < request.len) { + request[request_len] = try reader.interface.takeByte(); + request_len += 1; + if (std.mem.endsWith(u8, request[0..request_len], "\r\n\r\n")) break; + } else return error.TestRequestTooLarge; + const key = headerValue(request[0 .. request_len - 4], "sec-websocket-key") orelse return error.TestMissingWebSocketKey; + var accept_buffer: [28]u8 = undefined; + const accept = websocketAccept(key, &accept_buffer); + var write_buffer: [4096]u8 = undefined; + var writer = client_stream.writer(zio, &write_buffer); + try writer.interface.print( + "HTTP/1.1 101 Switching Protocols\r\nConnection: Upgrade\r\nUpgrade: websocket\r\nSec-WebSocket-Accept: {s}\r\n\r\n", + .{accept}, + ); + try writer.interface.flush(); + self.upgraded.store(true, .seq_cst); + while (!self.stopping.load(.seq_cst)) { + var sleep_io: std.Io.Threaded = .init_single_threaded; + sleep_io.io().sleep(.fromMilliseconds(1), .real) catch {}; + } + } +}; +const LoopbackMode = enum { + never_accept, + hang_after_upgrade, + reset_after_upgrade, + complete_then_hang_close, + binary_then_close, + ping_then_complete, + oversized_frame, +}; + +const LoopbackWebSocketFixture = struct { + io_backend: std.Io.Threaded = .init_single_threaded, + server: std.Io.net.Server, + mode: LoopbackMode, + thread: ?std.Thread = null, + stopping: std.atomic.Value(bool) = .init(false), + upgraded: std.atomic.Value(bool) = .init(false), + failure: ?anyerror = null, + + fn init(mode: LoopbackMode) !@This() { + var fixture: @This() = .{ .server = undefined, .mode = mode }; + var address = try std.Io.net.IpAddress.parse("127.0.0.1", 0); + fixture.server = try address.listen(fixture.io(), .{ .reuse_address = true }); + return fixture; + } + + fn io(self: *@This()) std.Io { + return self.io_backend.io(); + } + + fn endpoint(self: *@This(), buffer: []u8) ![]const u8 { + return std.fmt.bufPrint(buffer, "http://127.0.0.1:{d}/responses", .{self.server.socket.address.getPort()}); + } + + fn start(self: *@This()) !void { + self.thread = try std.Thread.spawn(.{}, run, .{self}); + } + + fn deinit(self: *@This()) void { + self.stopping.store(true, .seq_cst); + if (self.thread) |thread| { + const listener = std.Io.net.Stream{ .socket = self.server.socket }; + listener.shutdown(self.io(), .both) catch {}; + thread.join(); + self.thread = null; + } + self.server.deinit(self.io()); + } + + fn hold(self: *@This()) void { + while (!self.stopping.load(.seq_cst)) { + self.io().sleep(.fromMilliseconds(1), .real) catch {}; + } + } + + fn run(self: *@This()) void { + self.runFallible() catch |err| { + if (!self.stopping.load(.seq_cst)) self.failure = err; + }; + } + + fn runFallible(self: *@This()) !void { + if (self.mode == .never_accept) return self.hold(); + const zio = self.io(); + var client_stream = try self.server.accept(zio); + defer client_stream.close(zio); + if (self.stopping.load(.seq_cst)) return; + + var socket_buffer: [4096]u8 = undefined; + var reader = client_stream.reader(zio, &socket_buffer); + var request: [16 * 1024]u8 = undefined; + var request_len: usize = 0; + while (request_len < request.len) { + request[request_len] = try reader.interface.takeByte(); + request_len += 1; + if (std.mem.endsWith(u8, request[0..request_len], "\r\n\r\n")) break; + } else return error.TestRequestTooLarge; + const key = headerValue(request[0 .. request_len - 4], "sec-websocket-key") orelse return error.TestMissingWebSocketKey; + var accept_buffer: [28]u8 = undefined; + const accept = websocketAccept(key, &accept_buffer); + var write_buffer: [4096]u8 = undefined; + var writer = client_stream.writer(zio, &write_buffer); + try writer.interface.print( + "HTTP/1.1 101 Switching Protocols\r\nConnection: Upgrade\r\nUpgrade: websocket\r\nSec-WebSocket-Accept: {s}\r\n\r\n", + .{accept}, + ); + try writer.interface.flush(); + self.upgraded.store(true, .seq_cst); + + if (self.mode == .reset_after_upgrade) { + const reset_on_close: std.posix.linger = .{ .onoff = 1, .linger = 0 }; + try std.posix.setsockopt( + client_stream.socket.handle, + std.posix.SOL.SOCKET, + std.posix.SO.LINGER, + std.mem.asBytes(&reset_on_close), + ); + return; + } + if (self.mode == .hang_after_upgrade) return self.hold(); + + try discardClientFrame(&reader.interface); + switch (self.mode) { + .complete_then_hang_close => { + try writeServerFrame(&writer.interface, .text, "{\"type\":\"response.output_text.delta\",\"delta\":\"ok\"}"); + try writeServerFrame(&writer.interface, .text, "{\"type\":\"response.completed\",\"response\":{\"id\":\"r1\",\"status\":\"completed\"}}"); + try writer.interface.flush(); + self.hold(); + }, + .binary_then_close => { + try writeServerFrame(&writer.interface, .binary, &.{0}); + try writer.interface.flush(); + }, + .ping_then_complete => { + try writeServerFrame(&writer.interface, .ping, "hi"); + try writeServerFrame(&writer.interface, .text, "{\"type\":\"response.output_text.delta\",\"delta\":\"ok\"}"); + try writeServerFrame(&writer.interface, .text, "{\"type\":\"response.completed\",\"response\":{\"id\":\"r1\",\"status\":\"completed\"}}"); + try writeServerFrame(&writer.interface, .close, &.{ 0x03, 0xe8 }); + try writer.interface.flush(); + try discardClientFrame(&reader.interface); + }, + .oversized_frame => { + try writer.interface.writeAll(&.{ 0x81, 127 }); + try writer.interface.writeInt(u64, max_frame_bytes + 1, .big); + try writer.interface.flush(); + }, + else => unreachable, + } + } +}; + +fn writeServerFrame(writer: *std.Io.Writer, opcode: Opcode, payload: []const u8) !void { + try writer.writeByte(0x80 | @as(u8, @intFromEnum(opcode))); + if (payload.len < 126) { + try writer.writeByte(@intCast(payload.len)); + } else if (payload.len <= std.math.maxInt(u16)) { + try writer.writeByte(126); + try writer.writeInt(u16, @intCast(payload.len), .big); + } else { + try writer.writeByte(127); + try writer.writeInt(u64, @intCast(payload.len), .big); + } + try writer.writeAll(payload); +} + +fn discardClientFrame(reader: *std.Io.Reader) !void { + _ = try reader.takeByte(); + const second = try reader.takeByte(); + if (second & 0x80 == 0) return error.WebSocketProtocolViolation; + var length: u64 = second & 0x7f; + if (length == 126) length = try reader.takeInt(u16, .big) else if (length == 127) length = try reader.takeInt(u64, .big); + var mask: [4]u8 = undefined; + try reader.readSliceAll(&mask); + var discarded: [4096]u8 = undefined; + var remaining = length; + while (remaining > 0) { + const chunk_len: usize = @intCast(@min(remaining, discarded.len)); + try reader.readSliceAll(discarded[0..chunk_len]); + remaining -= chunk_len; + } +} + +fn headerValue(headers: []const u8, name: []const u8) ?[]const u8 { + var lines = std.mem.splitSequence(u8, headers, "\r\n"); + _ = lines.next(); + while (lines.next()) |line| { + const colon = std.mem.findScalar(u8, line, ':') orelse continue; + if (std.ascii.eqlIgnoreCase(std.mem.trim(u8, line[0..colon], " \t"), name)) { + return std.mem.trim(u8, line[colon + 1 ..], " \t"); + } + } + return null; +} + +test "WebSocket cancellation interrupts a backpressured response.create write" { + var fixture = try StalledWriteFixture.init(); + defer fixture.deinit(); + try fixture.start(); + + var endpoint_buffer: [128]u8 = undefined; + const payload = try std.testing.allocator.alloc(u8, 4 * 1024 * 1024); + defer std.testing.allocator.free(payload); + @memset(payload, 'x'); + var cancelled = std.atomic.Value(bool).init(false); + var delivery = gateway_client.DeliveryCertainty.init(); + const Canceller = struct { + fn run(server: *StalledWriteFixture, flag: *std.atomic.Value(bool)) void { + while (!server.upgraded.load(.seq_cst)) { + var sleep_io: std.Io.Threaded = .init_single_threaded; + sleep_io.io().sleep(.fromMilliseconds(1), .real) catch {}; + } + flag.store(true, .seq_cst); + } + }; + const canceller = try std.Thread.spawn(.{}, Canceller.run, .{ &fixture, &cancelled }); + defer canceller.join(); + const result = stream(std.testing.allocator, .{ + .endpoint = try fixture.endpoint(&endpoint_buffer), + .authorization = "Bearer test", + .account_id = "test", + .session_id = null, + .payload = payload, + .deadline = null, + .cancel_flag = &cancelled, + .delivery = &delivery, + }, @ptrCast(&cancelled), struct { + fn ignore(_: *anyopaque, _: []const u8) !bool { + return false; + } + }.ignore); + try std.testing.expectError(error.Cancelled, result); + try std.testing.expectEqual(gateway_client.DeliveryCertainty.State.possibly_sent, delivery.load()); + if (fixture.failure) |err| return err; +} + +test "WebSocket cancellation interrupts a stalled connect" { + var cancelled = std.atomic.Value(bool).init(false); + var delivery = gateway_client.DeliveryCertainty.init(); + const Canceller = struct { + fn run(flag: *std.atomic.Value(bool)) void { + io_mod.sleep(20 * std.time.ns_per_ms); + flag.store(true, .seq_cst); + } + }; + const canceller = try std.Thread.spawn(.{}, Canceller.run, .{&cancelled}); + const started = std.Io.Clock.Timestamp.now(io_mod.getIo(), .awake); + const result = stream(std.testing.allocator, .{ + .endpoint = "http://192.0.2.1:9/responses", + .authorization = "Bearer test", + .account_id = "test", + .session_id = null, + .payload = "{}", + .deadline = null, + .cancel_flag = &cancelled, + .delivery = &delivery, + }, @ptrCast(&cancelled), struct { + fn ignore(_: *anyopaque, _: []const u8) !bool { + return false; + } + }.ignore); + const elapsed_ms = started.durationTo(std.Io.Clock.Timestamp.now(io_mod.getIo(), .awake)).raw.toMilliseconds(); + canceller.join(); + if (result) |_| { + return error.TestExpectedError; + } else |err| { + if (err == error.Cancelled) { + try std.testing.expect(elapsed_ms < 2_000); + try std.testing.expectEqual(gateway_client.DeliveryCertainty.State.definitely_unsent, delivery.load()); + return; + } + if (elapsed_ms >= 5) return err; + } + + var fixture = try LoopbackWebSocketFixture.init(.never_accept); + defer fixture.deinit(); + try fixture.start(); + var endpoint_buffer: [128]u8 = undefined; + cancelled.store(false, .seq_cst); + delivery = gateway_client.DeliveryCertainty.init(); + const fallback_canceller = try std.Thread.spawn(.{}, Canceller.run, .{&cancelled}); + const fallback_started = std.Io.Clock.Timestamp.now(io_mod.getIo(), .awake); + const fallback_result = stream(std.testing.allocator, .{ + .endpoint = try fixture.endpoint(&endpoint_buffer), + .authorization = "Bearer test", + .account_id = "test", + .session_id = null, + .payload = "{}", + .deadline = null, + .cancel_flag = &cancelled, + .delivery = &delivery, + }, @ptrCast(&cancelled), struct { + fn ignore(_: *anyopaque, _: []const u8) !bool { + return false; + } + }.ignore); + const fallback_elapsed_ms = fallback_started.durationTo(std.Io.Clock.Timestamp.now(io_mod.getIo(), .awake)).raw.toMilliseconds(); + fallback_canceller.join(); + try std.testing.expectError(error.Cancelled, fallback_result); + try std.testing.expect(fallback_elapsed_ms < 2_000); + try std.testing.expectEqual(gateway_client.DeliveryCertainty.State.definitely_unsent, delivery.load()); + if (fixture.failure) |err| return err; +} + +test "WebSocket cancellation interrupts a hung close handshake" { + var fixture = try LoopbackWebSocketFixture.init(.complete_then_hang_close); + defer fixture.deinit(); + try fixture.start(); + var endpoint_buffer: [128]u8 = undefined; + var cancelled = std.atomic.Value(bool).init(false); + var delivery = gateway_client.DeliveryCertainty.init(); + const Canceller = struct { + fn run(server: *LoopbackWebSocketFixture, flag: *std.atomic.Value(bool)) void { + while (!server.upgraded.load(.seq_cst)) io_mod.sleep(std.time.ns_per_ms); + io_mod.sleep(50 * std.time.ns_per_ms); + flag.store(true, .seq_cst); + } + }; + const canceller = try std.Thread.spawn(.{}, Canceller.run, .{ &fixture, &cancelled }); + defer canceller.join(); + const result = stream(std.testing.allocator, .{ + .endpoint = try fixture.endpoint(&endpoint_buffer), + .authorization = "Bearer test", + .account_id = "test", + .session_id = null, + .payload = "{}", + .deadline = null, + .cancel_flag = &cancelled, + .delivery = &delivery, + }, @ptrCast(&cancelled), struct { + fn completed(_: *anyopaque, json: []const u8) !bool { + return std.mem.find(u8, json, "\"response.completed\"") != null; + } + }.completed); + try std.testing.expectError(error.Cancelled, result); + try std.testing.expectEqual(gateway_client.DeliveryCertainty.State.possibly_sent, delivery.load()); + if (fixture.failure) |err| return err; +} + +test "WebSocket peer reset after upgrade leaves delivery possibly sent" { + var fixture = try LoopbackWebSocketFixture.init(.reset_after_upgrade); + defer fixture.deinit(); + try fixture.start(); + var endpoint_buffer: [128]u8 = undefined; + const payload = try std.testing.allocator.alloc(u8, 4 * 1024); + defer std.testing.allocator.free(payload); + @memset(payload, 'x'); + var cancelled = std.atomic.Value(bool).init(false); + var delivery = gateway_client.DeliveryCertainty.init(); + const result = stream(std.testing.allocator, .{ + .endpoint = try fixture.endpoint(&endpoint_buffer), + .authorization = "Bearer test", + .account_id = "test", + .session_id = null, + .payload = payload, + .deadline = null, + .cancel_flag = &cancelled, + .delivery = &delivery, + }, @ptrCast(&cancelled), struct { + fn ignore(_: *anyopaque, _: []const u8) !bool { + return false; + } + }.ignore); + if (result) |_| return error.TestExpectedError else |err| try std.testing.expect(err != error.Cancelled); + try std.testing.expectEqual(gateway_client.DeliveryCertainty.State.possibly_sent, delivery.load()); + if (fixture.failure) |err| return err; +} + +test "WebSocket rejects unexpected binary frames" { + var fixture = try LoopbackWebSocketFixture.init(.binary_then_close); + defer fixture.deinit(); + try fixture.start(); + var endpoint_buffer: [128]u8 = undefined; + var cancelled = std.atomic.Value(bool).init(false); + var delivery = gateway_client.DeliveryCertainty.init(); + const result = stream(std.testing.allocator, .{ + .endpoint = try fixture.endpoint(&endpoint_buffer), + .authorization = "Bearer test", + .account_id = "test", + .session_id = null, + .payload = "{}", + .deadline = null, + .cancel_flag = &cancelled, + .delivery = &delivery, + }, @ptrCast(&cancelled), struct { + fn ignore(_: *anyopaque, _: []const u8) !bool { + return false; + } + }.ignore); + try std.testing.expectError(error.WebSocketUnexpectedBinary, result); + try std.testing.expectEqual(gateway_client.DeliveryCertainty.State.possibly_sent, delivery.load()); + if (fixture.failure) |err| return err; +} + +test "WebSocket answers ping then completes" { + var fixture = try LoopbackWebSocketFixture.init(.ping_then_complete); + defer fixture.deinit(); + try fixture.start(); + var endpoint_buffer: [128]u8 = undefined; + var cancelled = std.atomic.Value(bool).init(false); + var delivery = gateway_client.DeliveryCertainty.init(); + var completed = false; + try stream(std.testing.allocator, .{ + .endpoint = try fixture.endpoint(&endpoint_buffer), + .authorization = "Bearer test", + .account_id = "test", + .session_id = null, + .payload = "{}", + .deadline = null, + .cancel_flag = &cancelled, + .delivery = &delivery, + }, @ptrCast(&completed), struct { + fn handle(context: *anyopaque, json: []const u8) !bool { + const done: *bool = @ptrCast(@alignCast(context)); + if (std.mem.find(u8, json, "\"response.completed\"") != null) { + done.* = true; + return true; + } + return false; + } + }.handle); + try std.testing.expect(completed); + if (fixture.failure) |err| return err; +} + +test "WebSocket rejects inbound frames over max_frame_bytes" { + var fixture = try LoopbackWebSocketFixture.init(.oversized_frame); + defer fixture.deinit(); + try fixture.start(); + var endpoint_buffer: [128]u8 = undefined; + var cancelled = std.atomic.Value(bool).init(false); + var delivery = gateway_client.DeliveryCertainty.init(); + const result = stream(std.testing.allocator, .{ + .endpoint = try fixture.endpoint(&endpoint_buffer), + .authorization = "Bearer test", + .account_id = "test", + .session_id = null, + .payload = "{}", + .deadline = null, + .cancel_flag = &cancelled, + .delivery = &delivery, + }, @ptrCast(&cancelled), struct { + fn ignore(_: *anyopaque, _: []const u8) !bool { + return false; + } + }.ignore); + try std.testing.expectError(error.WebSocketMessageTooLarge, result); + if (fixture.failure) |err| return err; +} + +test "appendMessage rejects a 64 MiB overflow" { + var message: std.ArrayList(u8) = .empty; + defer message.deinit(std.testing.allocator); + const payload = try std.testing.allocator.alloc(u8, max_message_bytes); + defer std.testing.allocator.free(payload); + try appendMessage(&message, std.testing.allocator, payload); + try std.testing.expectError(error.WebSocketMessageTooLarge, appendMessage(&message, std.testing.allocator, "x")); +} + +test "writeFrame rejects payloads over max_message_bytes" { + var encoded: std.Io.Writer.Allocating = .init(std.testing.allocator); + defer encoded.deinit(); + const payload = try std.testing.allocator.alloc(u8, max_message_bytes + 1); + defer std.testing.allocator.free(payload); + try std.testing.expectError(error.WebSocketMessageTooLarge, writeFrame(&encoded.writer, .text, payload)); +} diff --git a/src/main.zig b/src/main.zig index bfcedb050..a68c21b99 100644 --- a/src/main.zig +++ b/src/main.zig @@ -70,6 +70,7 @@ const js_host_workspace = @import("core/hosts/js_host_workspace.zig"); const host_target = @import("core/hosts/target.zig"); const native_host = @import("core/hosts/native.zig"); const debug_trace = @import("core/shared/debug_trace.zig"); +const openai_codex = @import("gateway/openai_codex.zig"); const display_width = @import("core/shared/display_width.zig"); const file_index_mod = @import("core/workspace/file_index.zig"); const mcp_command_provider = @import("core/mcp/command_provider.zig"); @@ -3064,7 +3065,10 @@ fn mainC(c_argc: c_int, c_argv: [*][*:0]c_char, c_envp: [*:null]?[*:0]c_char) !v }); defer threaded.deinit(); io_mod.setIo(threaded.io()); - defer debug_trace.shutdown(); + defer { + debug_trace.shutdown(); + openai_codex.shutdownWebSockets(); + } debug_trace.configureFromEnv(processAllocator(), "."); try terminal_host.run( processAllocator(), @@ -3145,6 +3149,7 @@ fn runNonBenchmark(raw_args: []const [*:0]const u8, raw_env: RawEnviron, cli_arg }); if (early_threaded) |*threaded| io_mod.setIo(threaded.io()); } + defer openai_codex.shutdownWebSockets(); const before = try app_entry_runtime.runBeforeInteractive(alloc, cli_args, cfg); switch (before) { @@ -3160,7 +3165,10 @@ fn runNonBenchmark(raw_args: []const [*:0]const u8, raw_env: RawEnviron, cli_arg var owned_launch = launch; defer owned_launch.deinit(alloc); - defer debug_trace.shutdown(); + defer { + debug_trace.shutdown(); + openai_codex.shutdownWebSockets(); + } const outcome = try app_entry_runtime.runInteractive(App, alloc, &owned_launch); switch (outcome) { @@ -3646,6 +3654,7 @@ test "session reset traces and clears active paste state" { try std.testing.expectEqual(@as(usize, 0), app.input_runtime.paste.decision_bytes); try std.testing.expectEqual(@as(usize, 0), app.input_runtime.edit_state.input.items.len); debug_trace.shutdown(); + openai_codex.shutdownWebSockets(); var trace_file = try std.Io.Dir.openFileAbsolute(io_mod.getIo(), trace_path, .{}); defer trace_file.close(io_mod.getIo()); diff --git a/tests/e2e/tui-auth-source-selection.test.ts b/tests/e2e/tui-auth-source-selection.test.ts index bac87d5b9..0ebbac434 100644 --- a/tests/e2e/tui-auth-source-selection.test.ts +++ b/tests/e2e/tui-auth-source-selection.test.ts @@ -837,6 +837,154 @@ function startFakeCodexToolLoop(options: { }; } +function startFakeCodexWebSocket(options: { + holdOpen?: boolean; + closeOnOpen?: number; + closeAfterMessage?: number; + closeAfterFirstMessage?: number; + closeAfterCompletion?: boolean; + toolThenClose?: boolean; + stallUpgrade?: boolean; + rejectPreviousOnce?: boolean; + toolRoundTrip?: boolean; + reasoningState?: boolean; + rejectUpgradeWithSse?: boolean; +} = {}) { + const requests: string[] = []; + const closeCodes: number[] = []; + let upgradeRequests = 0; + let sseRequests = 0; + let httpResponseRequests = 0; + let rejectedPrevious = false; + const accessToken = chatgptAccessToken("acct_websocket"); + const server = Bun.serve<{ opened: boolean }>({ + hostname: "127.0.0.1", + port: 0, + async fetch(request, server) { + const path = new URL(request.url).pathname; + if (path === "/models") { + return Response.json({ models: [ + { slug: "gpt-5.6-sol", visibility: "list", supported_in_api: true, supported_reasoning_levels: [{ effort: "high" }], additional_speed_tiers: [], input_modalities: ["text"], context_window: 272000 }, + ] }); + } + if (path === "/responses") { + if (request.headers.get("upgrade")?.toLowerCase() === "websocket") { + upgradeRequests += 1; + if (options.rejectUpgradeWithSse) return new Response("upgrade unavailable", { status: 426 }); + if (options.stallUpgrade) return new Promise(() => {}); + if (server.upgrade(request, { data: { opened: true } })) return; + } else { + httpResponseRequests += 1; + if (options.rejectUpgradeWithSse) { + requests.push(await request.text()); + sseRequests += 1; + const completed = { + type: "response.completed", + response: { + id: `resp_sse_${sseRequests}`, + status: "completed", + usage: { input_tokens: 5, output_tokens: 2 }, + }, + }; + return new Response( + `data: ${JSON.stringify({ type: "response.output_text.delta", delta: `CODEX_SSE_FALLBACK_${sseRequests}` })}\n\n` + + `data: ${JSON.stringify(completed)}\n\n`, + { headers: { "content-type": "text/event-stream" } }, + ); + } + } + } + return new Response("not found", { status: 404 }); + }, + websocket: { + open(ws) { + if (options.closeOnOpen !== undefined) ws.close(options.closeOnOpen, "fixture close"); + }, + message(ws, message) { + const payload = String(message); + requests.push(payload); + const parsed = JSON.parse(payload) as { previous_response_id?: string }; + if (options.rejectPreviousOnce && parsed.previous_response_id && !rejectedPrevious) { + rejectedPrevious = true; + ws.send(JSON.stringify({ + type: "error", + error: { code: "previous_response_not_found", message: "Previous response was not found." }, + })); + return; + } + if (options.toolRoundTrip && requests.length === 1) { + ws.send(JSON.stringify({ + type: "response.output_item.added", + output_index: 0, + item: { type: "function_call", call_id: "call_phase3", name: "read_file" }, + })); + ws.send(JSON.stringify({ + type: "response.function_call_arguments.done", + output_index: 0, + arguments: '{"path":"README.md"}', + })); + ws.send(JSON.stringify({ + type: "response.completed", + response: { id: "resp_websocket_1", status: "completed", usage: { input_tokens: 5, output_tokens: 2 } }, + })); + return; + } + if (options.toolThenClose) { + ws.send(JSON.stringify({ + type: "response.output_item.added", + output_index: 0, + item: { type: "function_call", call_id: "call_1", name: "read_file" }, + })); + ws.send(JSON.stringify({ + type: "response.function_call_arguments.delta", + output_index: 0, + delta: '{"path":"README.md"}', + })); + ws.close(1011, "fixture close"); + return; + } + if (options.closeAfterMessage !== undefined) { + ws.close(options.closeAfterMessage, "fixture close"); + return; + } + if (options.closeAfterFirstMessage !== undefined && requests.length === 1) { + ws.close(options.closeAfterFirstMessage, "fixture first-message close"); + return; + } + if (options.holdOpen) return; + if (options.reasoningState) { + ws.send(JSON.stringify({ + type: "response.output_item.done", + output_index: 0, + item: { type: "reasoning", id: `reasoning_${requests.length}`, encrypted_content: "opaque" }, + })); + } + ws.send(JSON.stringify({ type: "response.output_text.delta", delta: `CODEX_WEBSOCKET_OK_${requests.length}` })); + ws.send(JSON.stringify({ + type: "response.completed", + response: { id: `resp_websocket_${requests.length}`, status: "completed", usage: { input_tokens: 5, output_tokens: 2 } }, + })); + if (options.closeAfterCompletion) ws.close(1000, "fixture completed"); + }, + close(_ws, code) { + closeCodes.push(code); + }, + }, + }); + return { + accessToken, + requests, + closeCodes, + get upgradeRequests() { return upgradeRequests; }, + get sseRequests() { return sseRequests; }, + get httpResponseRequests() { return httpResponseRequests; }, + responsesUrl: `http://127.0.0.1:${server.port}/responses`, + modelsUrl: `http://127.0.0.1:${server.port}/models`, + stop() { server.stop(true); }, + }; +} + + function startFakeCodexCapacityLoop() { const bodies: string[] = []; const accessToken = chatgptAccessToken("acct_capacity_loop"); @@ -2762,6 +2910,592 @@ tmuxTest( 60_000, ); +test( + "Codex WebSocket streams a completion through the freshly built binary", + async () => { + home = mkdtempSync(join(tmpdir(), "fx-codex-websocket-")); + gateway = startFakeGateway([]); + const codex = startFakeCodexWebSocket({ reasoningState: true }); + try { + writeSeededChatGptLogin(home, codex.accessToken); + writeFileSync( + join(home, ".fx", "settings.json"), + JSON.stringify({ provider: "codex", codex_model: "gpt-5.6-sol" }) + "\n", + { mode: 0o600 }, + ); + const result = await runFx( + ["ask", "--json", "--auto", "--no-save", "Use the WebSocket transport."], + { + env: { + HOME: home, + AI_GATEWAY_API_KEY: "gateway-websocket-sentinel", + VERCEL_OIDC_TOKEN: undefined, + FX_DISABLE_KEYCHAIN: "1", + FX_AUTO_UPGRADE: "0", + FX_CODEX_TRANSPORT: "websocket", + FX_GATEWAY_BASE_URL: gateway.baseUrl, + FX_E2E_GATEWAY_MODELS_URL: `${gateway.baseUrl}/coding-agent/v1/models`, + FX_E2E_OPENAI_CODEX_RESPONSES_URL: codex.responsesUrl, + FX_E2E_OPENAI_CODEX_MODELS_URL: codex.modelsUrl, + }, + timeoutMs: TIMEOUT, + }, + ); + expect(result.code, `stdout: ${result.stdout}\nstderr: ${result.stderr}`).toBe(0); + expect(result.stdout).toContain("CODEX_WEBSOCKET_OK"); + expect(codex.requests).toHaveLength(1); + expect(codex.requests[0]).toContain('"type":"response.create"'); + expect(codex.requests[0]).not.toContain('"stream"'); + expect(gateway.requests).toHaveLength(0); + } finally { + codex.stop(); + } + }, + 60_000, +); + +test( + "Codex WebSocket policy close is not retried as a transport failure", + async () => { + home = mkdtempSync(join(tmpdir(), "fx-codex-websocket-policy-close-")); + gateway = startFakeGateway([]); + const codex = startFakeCodexWebSocket({ closeOnOpen: 1008 }); + try { + writeSeededChatGptLogin(home, codex.accessToken); + writeFileSync( + join(home, ".fx", "settings.json"), + JSON.stringify({ provider: "codex", codex_model: "gpt-5.6-sol" }) + "\n", + { mode: 0o600 }, + ); + const result = await runFx( + ["ask", "--json", "--auto", "--no-save", "Reject this request by policy."], + { + env: { + HOME: home, + AI_GATEWAY_API_KEY: "gateway-websocket-policy-sentinel", + VERCEL_OIDC_TOKEN: undefined, + FX_DISABLE_KEYCHAIN: "1", + FX_AUTO_UPGRADE: "0", + FX_CODEX_TRANSPORT: "websocket", + FX_GATEWAY_BASE_URL: gateway.baseUrl, + FX_E2E_GATEWAY_MODELS_URL: `${gateway.baseUrl}/coding-agent/v1/models`, + FX_E2E_OPENAI_CODEX_RESPONSES_URL: codex.responsesUrl, + FX_E2E_OPENAI_CODEX_MODELS_URL: codex.modelsUrl, + }, + timeoutMs: TIMEOUT, + }, + ); + expect(result.code).toBe(1); + expect(`${result.stdout}\n${result.stderr}`).toContain("WebSocketPolicyClosed"); + expect(codex.requests).toHaveLength(0); + expect(codex.upgradeRequests).toBe(1); + expect(gateway.requests).toHaveLength(0); + } finally { + codex.stop(); + } + }, + 60_000, +); + +test( + "Codex WebSocket continues a completed tool call with only its result", + async () => { + home = mkdtempSync(join(tmpdir(), "fx-codex-websocket-tool-continuation-")); + gateway = startFakeGateway([]); + const codex = startFakeCodexWebSocket({ toolRoundTrip: true }); + try { + writeSeededChatGptLogin(home, codex.accessToken); + writeFileSync( + join(home, ".fx", "settings.json"), + JSON.stringify({ provider: "codex", codex_model: "gpt-5.6-sol", permission_mode: "yolo" }) + "\n", + { mode: 0o600 }, + ); + const result = await runFx( + ["ask", "--json", "--auto", "--no-save", "Read README.md, then report success."], + { + env: { + HOME: home, + AI_GATEWAY_API_KEY: "gateway-websocket-tool-continuation-sentinel", + VERCEL_OIDC_TOKEN: undefined, + FX_DISABLE_KEYCHAIN: "1", + FX_AUTO_UPGRADE: "0", + FX_CODEX_TRANSPORT: "websocket", + FX_GATEWAY_BASE_URL: gateway.baseUrl, + FX_E2E_GATEWAY_MODELS_URL: `${gateway.baseUrl}/coding-agent/v1/models`, + FX_E2E_OPENAI_CODEX_RESPONSES_URL: codex.responsesUrl, + FX_E2E_OPENAI_CODEX_MODELS_URL: codex.modelsUrl, + }, + timeoutMs: TIMEOUT, + }, + ); + expect(result.code, `stdout: ${result.stdout}\nstderr: ${result.stderr}`).toBe(0); + expect(codex.upgradeRequests).toBe(1); + expect(codex.requests).toHaveLength(2); + expect(codex.requests[1]).toContain('"previous_response_id":"resp_websocket_1"'); + expect(codex.requests[1]).toContain('"type":"function_call_output"'); + expect(codex.requests[1]).not.toContain("Read README.md, then report success."); + expect(gateway.requests).toHaveLength(0); + } finally { + codex.stop(); + } + }, + 60_000, +); + +test( + "Codex WebSocket close after response.create never replays the turn", + async () => { + home = mkdtempSync(join(tmpdir(), "fx-codex-websocket-post-send-close-")); + gateway = startFakeGateway([]); + const codex = startFakeCodexWebSocket({ closeAfterMessage: 1011 }); + try { + writeSeededChatGptLogin(home, codex.accessToken); + writeFileSync( + join(home, ".fx", "settings.json"), + JSON.stringify({ provider: "codex", codex_model: "gpt-5.6-sol" }) + "\n", + { mode: 0o600 }, + ); + const result = await runFx( + ["ask", "--json", "--auto", "--no-save", "Do not replay this request."], + { + env: { + HOME: home, + AI_GATEWAY_API_KEY: "gateway-websocket-post-send-close-sentinel", + VERCEL_OIDC_TOKEN: undefined, + FX_DISABLE_KEYCHAIN: "1", + FX_AUTO_UPGRADE: "0", + FX_CODEX_TRANSPORT: "websocket", + FX_GATEWAY_BASE_URL: gateway.baseUrl, + FX_E2E_GATEWAY_MODELS_URL: `${gateway.baseUrl}/coding-agent/v1/models`, + FX_E2E_OPENAI_CODEX_RESPONSES_URL: codex.responsesUrl, + FX_E2E_OPENAI_CODEX_MODELS_URL: codex.modelsUrl, + }, + timeoutMs: TIMEOUT, + }, + ); + expect(result.code).toBe(1); + expect(`${result.stdout}\n${result.stderr}`).toContain("WebSocketClosedBeforeCompletion"); + expect(codex.requests).toHaveLength(1); + expect(codex.upgradeRequests).toBe(1); + expect(codex.httpResponseRequests).toBe(0); + expect(gateway.requests).toHaveLength(0); + } finally { + codex.stop(); + } + }, + 60_000, +); + +tmuxTest( + "Codex WebSocket uses full context when persisted tool evidence changes history", + async () => { + home = mkdtempSync(join(tmpdir(), "fx-codex-websocket-tool-turn-continuation-")); + stderrPath = join(home, "stderr.log"); + writeFileSync(stderrPath, ""); + gateway = startFakeGateway([]); + const codex = startFakeCodexWebSocket({ toolRoundTrip: true, reasoningState: true }); + try { + writeSeededChatGptLogin(home, codex.accessToken); + writeFileSync( + join(home, ".fx", "settings.json"), + JSON.stringify({ provider: "codex", codex_model: "gpt-5.6-sol", permission_mode: "yolo" }) + "\n", + { mode: 0o600 }, + ); + session = await startFx(home, stderrPath, gateway, undefined, undefined, { + FX_MODEL: undefined, + FX_CODEX_TRANSPORT: "websocket", + FX_E2E_OPENAI_CODEX_RESPONSES_URL: codex.responsesUrl, + FX_E2E_OPENAI_CODEX_MODELS_URL: codex.modelsUrl, + }); + await session.waitForComposer(TIMEOUT); + await session.sendText("Read README.md for the first Phase 3 turn."); + await session.waitForText("CODEX_WEBSOCKET_OK_2", TIMEOUT); + await session.sendText("Complete the next Phase 3 turn without replay."); + await session.waitForText("CODEX_WEBSOCKET_OK_3", TIMEOUT); + + expect(codex.upgradeRequests).toBe(1); + expect(codex.requests).toHaveLength(3); + expect(codex.requests[1]).toContain('"previous_response_id":"resp_websocket_1"'); + expect(codex.requests[1]).toContain('"type":"function_call_output"'); + expect(codex.requests[2]).not.toContain("previous_response_id"); + expect(codex.requests[2]).toContain("Complete the next Phase 3 turn without replay."); + expect(codex.requests[2]).toContain("Read README.md for the first Phase 3 turn."); + expect(codex.requests[2]).toContain('"type":"function_call_output"'); + expect(readFileSync(stderrPath, "utf8")).toBe(""); + } finally { + codex.stop(); + } + }, + 60_000, +); + + +test( + "Codex WebSocket tool stream close never replays the turn", + async () => { + home = mkdtempSync(join(tmpdir(), "fx-codex-websocket-tool-close-")); + gateway = startFakeGateway([]); + const codex = startFakeCodexWebSocket({ toolThenClose: true }); + try { + writeSeededChatGptLogin(home, codex.accessToken); + writeFileSync( + join(home, ".fx", "settings.json"), + JSON.stringify({ provider: "codex", codex_model: "gpt-5.6-sol" }) + "\n", + { mode: 0o600 }, + ); + const result = await runFx( + ["ask", "--json", "--auto", "--no-save", "Do not replay this tool request."], + { + env: { + HOME: home, + AI_GATEWAY_API_KEY: "gateway-websocket-tool-close-sentinel", + VERCEL_OIDC_TOKEN: undefined, + FX_DISABLE_KEYCHAIN: "1", + FX_AUTO_UPGRADE: "0", + FX_CODEX_TRANSPORT: "websocket", + FX_GATEWAY_BASE_URL: gateway.baseUrl, + FX_E2E_GATEWAY_MODELS_URL: `${gateway.baseUrl}/coding-agent/v1/models`, + FX_E2E_OPENAI_CODEX_RESPONSES_URL: codex.responsesUrl, + FX_E2E_OPENAI_CODEX_MODELS_URL: codex.modelsUrl, + }, + timeoutMs: TIMEOUT, + }, + ); + expect(result.code).toBe(1); + expect(`${result.stdout}\n${result.stderr}`).toContain("WebSocketClosedBeforeCompletion"); + expect(codex.requests).toHaveLength(1); + expect(codex.upgradeRequests).toBe(1); + expect(gateway.requests).toHaveLength(0); + } finally { + codex.stop(); + } + }, + 60_000, +); + +tmuxTest( + "Codex WebSocket reuses one socket for sequential interactive turns", + async () => { + home = mkdtempSync(join(tmpdir(), "fx-codex-websocket-reuse-")); + stderrPath = join(home, "stderr.log"); + writeFileSync(stderrPath, ""); + gateway = startFakeGateway([]); + const codex = startFakeCodexWebSocket({ reasoningState: true }); + try { + writeSeededChatGptLogin(home, codex.accessToken); + writeFileSync( + join(home, ".fx", "settings.json"), + JSON.stringify({ provider: "codex", codex_model: "gpt-5.6-sol" }) + "\n", + { mode: 0o600 }, + ); + session = await startFx(home, stderrPath, gateway, undefined, undefined, { + FX_MODEL: undefined, + FX_CODEX_TRANSPORT: "websocket", + FX_E2E_OPENAI_CODEX_RESPONSES_URL: codex.responsesUrl, + FX_E2E_OPENAI_CODEX_MODELS_URL: codex.modelsUrl, + }); + await session.waitForComposer(TIMEOUT); + await session.sendText("Complete the first retained-socket turn."); + await session.waitForText("CODEX_WEBSOCKET_OK_1", TIMEOUT); + await session.sendText("Complete the second retained-socket turn."); + const secondDeadline = Date.now() + TIMEOUT; + while (codex.requests.length < 2) { + if (Date.now() >= secondDeadline) throw new Error("Second Codex WebSocket request did not arrive"); + await Bun.sleep(25); + } + await session.waitForText("CODEX_WEBSOCKET_OK_2", TIMEOUT); + expect(codex.upgradeRequests).toBe(1); + expect(codex.requests).toHaveLength(2); + expect(codex.requests.every((request) => request.includes('"type":"response.create"'))).toBe(true); + expect(codex.requests[1]).toContain("Complete the second retained-socket turn."); + expect(codex.requests[1]).not.toContain("Complete the first retained-socket turn."); + expect(codex.requests[1]).not.toContain("CODEX_WEBSOCKET_OK_1"); + expect(codex.requests[1]).toContain('"previous_response_id":"resp_websocket_1"'); + expect(readFileSync(stderrPath, "utf8")).toBe(""); + } finally { + codex.stop(); + } + }, + 60_000, +); + +tmuxTest( + "Codex WebSocket handshake failure falls back once and latches SSE", + async () => { + home = mkdtempSync(join(tmpdir(), "fx-codex-websocket-fallback-")); + stderrPath = join(home, "stderr.log"); + writeFileSync(stderrPath, ""); + gateway = startFakeGateway([]); + const codex = startFakeCodexWebSocket({ rejectUpgradeWithSse: true }); + try { + writeSeededChatGptLogin(home, codex.accessToken); + writeFileSync( + join(home, ".fx", "settings.json"), + JSON.stringify({ provider: "codex", codex_model: "gpt-5.6-sol" }) + "\n", + { mode: 0o600 }, + ); + session = await startFx(home, stderrPath, gateway, undefined, undefined, { + FX_MODEL: undefined, + FX_CODEX_TRANSPORT: "websocket", + FX_E2E_OPENAI_CODEX_RESPONSES_URL: codex.responsesUrl, + FX_E2E_OPENAI_CODEX_MODELS_URL: codex.modelsUrl, + }); + await session.waitForComposer(TIMEOUT); + await session.sendText("Fall back after the rejected WebSocket upgrade."); + await session.waitForText("CODEX_SSE_FALLBACK_1", TIMEOUT); + await session.sendText("Keep using SSE after the fallback latch is armed."); + await session.waitForText("CODEX_SSE_FALLBACK_2", TIMEOUT); + expect(codex.upgradeRequests).toBe(1); + expect(codex.sseRequests).toBe(2); + expect(codex.requests).toHaveLength(2); + expect(readFileSync(stderrPath, "utf8")).toBe(""); + } finally { + codex.stop(); + } + }, + 60_000, +); + +tmuxTest( + "Codex WebSocket reconnects before delivery when the retained socket closes", + async () => { + home = mkdtempSync(join(tmpdir(), "fx-codex-websocket-reconnect-")); + stderrPath = join(home, "stderr.log"); + writeFileSync(stderrPath, ""); + gateway = startFakeGateway([]); + const codex = startFakeCodexWebSocket({ closeAfterCompletion: true }); + try { + writeSeededChatGptLogin(home, codex.accessToken); + writeFileSync( + join(home, ".fx", "settings.json"), + JSON.stringify({ provider: "codex", codex_model: "gpt-5.6-sol" }) + "\n", + { mode: 0o600 }, + ); + session = await startFx(home, stderrPath, gateway, undefined, undefined, { + FX_MODEL: undefined, + FX_CODEX_TRANSPORT: "websocket", + FX_E2E_OPENAI_CODEX_RESPONSES_URL: codex.responsesUrl, + FX_E2E_OPENAI_CODEX_MODELS_URL: codex.modelsUrl, + }); + await session.waitForComposer(TIMEOUT); + await session.sendText("Complete before the retained socket closes."); + await session.waitForText("CODEX_WEBSOCKET_OK_1", TIMEOUT); + await session.sendText("Reconnect without replaying either turn."); + await session.waitForText("CODEX_WEBSOCKET_OK_2", TIMEOUT); + expect(codex.upgradeRequests).toBe(2); + expect(codex.requests).toHaveLength(2); + expect(codex.requests[1]).toContain("Complete before the retained socket closes."); + expect(codex.requests[1]).toContain("Reconnect without replaying either turn."); + expect(readFileSync(stderrPath, "utf8")).toBe(""); + } finally { + codex.stop(); + } + }, + 60_000, +); + +tmuxTest( + "Codex WebSocket retries full context when continuation state is missing", + async () => { + home = mkdtempSync(join(tmpdir(), "fx-codex-websocket-continuation-recovery-")); + stderrPath = join(home, "stderr.log"); + writeFileSync(stderrPath, ""); + gateway = startFakeGateway([]); + const codex = startFakeCodexWebSocket({ rejectPreviousOnce: true }); + try { + writeSeededChatGptLogin(home, codex.accessToken); + writeFileSync( + join(home, ".fx", "settings.json"), + JSON.stringify({ provider: "codex", codex_model: "gpt-5.6-sol" }) + "\n", + { mode: 0o600 }, + ); + session = await startFx(home, stderrPath, gateway, undefined, undefined, { + FX_MODEL: undefined, + FX_CODEX_TRANSPORT: "websocket", + FX_E2E_OPENAI_CODEX_RESPONSES_URL: codex.responsesUrl, + FX_E2E_OPENAI_CODEX_MODELS_URL: codex.modelsUrl, + }); + await session.waitForComposer(TIMEOUT); + await session.sendText("Complete the first continuation-recovery turn."); + await session.waitForText("CODEX_WEBSOCKET_OK_1", TIMEOUT); + await session.sendText("Recover the second continuation-recovery turn."); + await session.waitForText("CODEX_WEBSOCKET_OK_3", TIMEOUT); + + expect(codex.upgradeRequests).toBe(2); + expect(codex.requests).toHaveLength(3); + expect(codex.requests[1]).toContain('"previous_response_id":"resp_websocket_1"'); + expect(codex.requests[1]).not.toContain("Complete the first continuation-recovery turn."); + expect(codex.requests[2]).not.toContain("previous_response_id"); + expect(codex.requests[2]).toContain("Complete the first continuation-recovery turn."); + expect(codex.requests[2]).toContain("CODEX_WEBSOCKET_OK_1"); + expect(codex.requests[2]).toContain("Recover the second continuation-recovery turn."); + expect(readFileSync(stderrPath, "utf8")).toBe(""); + } finally { + codex.stop(); + } + }, + 60_000, +); + +tmuxTest( + "Codex WebSocket poisons a failed stream and recovers on the next turn", + async () => { + home = mkdtempSync(join(tmpdir(), "fx-codex-websocket-poison-recovery-")); + stderrPath = join(home, "stderr.log"); + writeFileSync(stderrPath, ""); + gateway = startFakeGateway([]); + const codex = startFakeCodexWebSocket({ closeAfterFirstMessage: 1011 }); + try { + writeSeededChatGptLogin(home, codex.accessToken); + writeFileSync( + join(home, ".fx", "settings.json"), + JSON.stringify({ provider: "codex", codex_model: "gpt-5.6-sol" }) + "\n", + { mode: 0o600 }, + ); + session = await startFx(home, stderrPath, gateway, undefined, undefined, { + FX_MODEL: undefined, + FX_CODEX_TRANSPORT: "websocket", + FX_E2E_OPENAI_CODEX_RESPONSES_URL: codex.responsesUrl, + FX_E2E_OPENAI_CODEX_MODELS_URL: codex.modelsUrl, + }); + await session.waitForComposer(TIMEOUT); + await session.sendText("Fail this retained-socket turn once."); + await session.waitForText("WebSocketClosedBeforeCompletion", TIMEOUT); + await session.sendText("Recover on a fresh socket."); + await session.waitForText("CODEX_WEBSOCKET_OK_2", TIMEOUT); + expect(codex.upgradeRequests).toBe(2); + expect(codex.requests).toHaveLength(2); + expect(codex.requests[0]).toContain("Fail this retained-socket turn once."); + expect(codex.requests[1]).toContain("Recover on a fresh socket."); + expect(readFileSync(stderrPath, "utf8")).toBe(""); + } finally { + codex.stop(); + } + }, + 60_000, +); + +tmuxTest( + "Codex WebSocket evicts a retained socket after its configured maximum age", + async () => { + home = mkdtempSync(join(tmpdir(), "fx-codex-websocket-age-")); + stderrPath = join(home, "stderr.log"); + writeFileSync(stderrPath, ""); + gateway = startFakeGateway([]); + const codex = startFakeCodexWebSocket(); + try { + writeSeededChatGptLogin(home, codex.accessToken); + writeFileSync( + join(home, ".fx", "settings.json"), + JSON.stringify({ provider: "codex", codex_model: "gpt-5.6-sol" }) + "\n", + { mode: 0o600 }, + ); + session = await startFx(home, stderrPath, gateway, undefined, undefined, { + FX_MODEL: undefined, + FX_CODEX_TRANSPORT: "websocket", + FX_CODEX_WEBSOCKET_MAX_CONNECTION_AGE_MS: "1", + FX_E2E_OPENAI_CODEX_RESPONSES_URL: codex.responsesUrl, + FX_E2E_OPENAI_CODEX_MODELS_URL: codex.modelsUrl, + }); + await session.waitForComposer(TIMEOUT); + await session.sendText("Complete on the initial short-lived socket."); + await session.waitForText("CODEX_WEBSOCKET_OK_1", TIMEOUT); + await Bun.sleep(20); + await session.sendText("Complete after connection-age eviction."); + await session.waitForText("CODEX_WEBSOCKET_OK_2", TIMEOUT); + expect(codex.upgradeRequests).toBe(2); + expect(codex.requests).toHaveLength(2); + expect(readFileSync(stderrPath, "utf8")).toBe(""); + } finally { + codex.stop(); + } + }, + 60_000, +); + +tmuxTest( + "Codex WebSocket cancellation unblocks a stalled upgrade", + async () => { + home = mkdtempSync(join(tmpdir(), "fx-codex-websocket-upgrade-cancel-")); + stderrPath = join(home, "stderr.log"); + writeFileSync(stderrPath, ""); + gateway = startFakeGateway([]); + const codex = startFakeCodexWebSocket({ stallUpgrade: true }); + try { + writeSeededChatGptLogin(home, codex.accessToken); + writeFileSync( + join(home, ".fx", "settings.json"), + JSON.stringify({ provider: "codex", codex_model: "gpt-5.6-sol" }) + "\n", + { mode: 0o600 }, + ); + session = await startFx(home, stderrPath, gateway, undefined, undefined, { + FX_MODEL: undefined, + FX_CODEX_TRANSPORT: "websocket", + FX_E2E_OPENAI_CODEX_RESPONSES_URL: codex.responsesUrl, + FX_E2E_OPENAI_CODEX_MODELS_URL: codex.modelsUrl, + }); + await session.waitForComposer(TIMEOUT); + await session.sendText("Cancel a stalled WebSocket upgrade."); + const upgradeDeadline = Date.now() + TIMEOUT; + while (codex.upgradeRequests === 0) { + if (Date.now() >= upgradeDeadline) throw new Error("Codex WebSocket upgrade did not arrive"); + await Bun.sleep(25); + } + await session.sendKeys("C-c"); + await session.waitForComposer(TIMEOUT); + expect(session.isAlive()).toBe(true); + expect(codex.upgradeRequests).toBe(1); + expect(codex.requests).toHaveLength(0); + expect(readFileSync(stderrPath, "utf8")).toBe(""); + } finally { + codex.stop(); + } + }, + 60_000, +); + + +tmuxTest( + "Codex WebSocket cancellation unblocks an idle response", + async () => { + home = mkdtempSync(join(tmpdir(), "fx-codex-websocket-cancel-")); + stderrPath = join(home, "stderr.log"); + writeFileSync(stderrPath, ""); + gateway = startFakeGateway([]); + const codex = startFakeCodexWebSocket({ holdOpen: true }); + try { + writeSeededChatGptLogin(home, codex.accessToken); + writeFileSync( + join(home, ".fx", "settings.json"), + JSON.stringify({ provider: "codex", codex_model: "gpt-5.6-sol" }) + "\n", + { mode: 0o600 }, + ); + session = await startFx(home, stderrPath, gateway, undefined, undefined, { + FX_MODEL: undefined, + FX_CODEX_TRANSPORT: "websocket", + FX_E2E_OPENAI_CODEX_RESPONSES_URL: codex.responsesUrl, + FX_E2E_OPENAI_CODEX_MODELS_URL: codex.modelsUrl, + }); + await session.waitForComposer(TIMEOUT); + await session.sendText("Wait for a held WebSocket response."); + const requestDeadline = Date.now() + TIMEOUT; + while (codex.requests.length === 0) { + if (Date.now() >= requestDeadline) throw new Error("Codex WebSocket request did not arrive"); + await Bun.sleep(25); + } + await session.sendKeys("C-c"); + await session.waitForComposer(TIMEOUT); + expect(session.isAlive()).toBe(true); + expect(codex.requests).toHaveLength(1); + expect(codex.upgradeRequests).toBe(1); + expect(readFileSync(stderrPath, "utf8")).toBe(""); + } finally { + codex.stop(); + } + }, + 60_000, +); + test( "ChatGPT tool loops round-trip encrypted reasoning without Gateway leakage", async () => {