diff --git a/README.md b/README.md index 51e146995..ba9b5d290 100644 --- a/README.md +++ b/README.md @@ -81,6 +81,37 @@ sideband connection. See the [Realtime WebSocket guide](realtime.md) and the run [text, transcription, voice, and sideband examples](examples/realtime/README.md) for lifecycle, authentication, proxy, TLS, and custom transport details. +### Responses WebSockets + +The Responses API also supports a persistent, block-scoped WebSocket for +sequential and multiplexed response creation. Add the optional +`async-websocket` gem, then send `response.create` events through +`client.responses.connect`: + +```ruby +client.responses.connect do |connection| + connection.response.create( + model: "gpt-5.2", + input: "Say hello.", + stream_id: "turn_1" + ) + + connection.each do |event| + print(event.delta) if event.type == :"response.output_text.delta" + break if event.type == :"response.completed" + end +end +``` + +The connection is intentionally single-owner: callers serialize writes and +use one reader (`receive` or `each`). The SDK forwards response fields and +`stream_id` values without imposing additional client-side policy. Known +server events are decoded into generated models on a best-effort basis, while +newer event types remain observable as `UnknownServerEvent` values. The SDK +does not automatically reconnect or replay an ambiguous write. When the server +closes a connection, open a new connection and continue with +`previous_response_id` when the response was stored. + ### Pagination List methods in the OpenAI API are paginated. diff --git a/lib/openai.rb b/lib/openai.rb index 3cfccd9d3..512040fa0 100644 --- a/lib/openai.rb +++ b/lib/openai.rb @@ -1315,4 +1315,6 @@ require_relative "openai/helpers/streaming/chat_events" require_relative "openai/helpers/streaming/chat_completion_stream" require_relative "openai/streaming" +require_relative "openai/helpers/websocket" require_relative "openai/helpers/realtime" +require_relative "openai/helpers/responses_websocket" diff --git a/lib/openai/helpers/realtime/client_extension.rb b/lib/openai/helpers/realtime/client_extension.rb index 996e6666c..f2e21a623 100644 --- a/lib/openai/helpers/realtime/client_extension.rb +++ b/lib/openai/helpers/realtime/client_extension.rb @@ -5,6 +5,8 @@ module Helpers module Realtime # Realtime request integration kept outside the generated client implementation. module ClientExtension + include OpenAI::WebSocket::ClientRequest + MAX_CALL_CLEANUP_SECONDS = 5.0 private_constant :MAX_CALL_CLEANUP_SECONDS @@ -140,28 +142,20 @@ def realtime_connection_request(path:, query:, websocket_base_url: nil, options: # # @api private def with_realtime_connection_request(path:, query:, websocket_base_url: nil, options: nil) - request, deadline = build_realtime_connection_request( - path: path, - query: query, - websocket_base_url: websocket_base_url, - options: options - ) - handshake_completed = false - mark_handshake_completed = -> { handshake_completed = true } - yield(request, mark_handshake_completed) - rescue OpenAI::Errors::RealtimeConnectionError => e - raise if handshake_completed - raise unless e.http_status == 401 && @workload_identity_auth + build = lambda do |deadline| + build_realtime_connection_request( + path: path, + query: query, + websocket_base_url: websocket_base_url, + options: options, + deadline: deadline + ) + end - @workload_identity_auth.invalidate_token - refreshed, = build_realtime_connection_request( - path: path, - query: query, - websocket_base_url: websocket_base_url, - options: options, - deadline: deadline - ) - yield(refreshed, mark_handshake_completed) + with_websocket_connection_retry( + error_class: OpenAI::Errors::RealtimeConnectionError, + build: build + ) { |request, marker| yield(request, marker) } end private def build_realtime_connection_request( @@ -171,6 +165,28 @@ def with_realtime_connection_request(path:, query:, websocket_base_url: nil, opt options:, deadline: nil ) + build_shared_websocket_connection_request( + path: path, + query: query, + websocket_base_url: websocket_base_url, + options: options, + deadline: deadline, + validate: -> (value) { validate_realtime_websocket_request!(value) }, + invalid_base_url_message: "`websocket_base_url` must be an absolute HTTP or WebSocket URL " \ + "without credentials, query, or fragment", + malformed_base_url_message: "`websocket_base_url` is not a valid URL", + preserve_base_url_cause: true, + extra_query_message: "`request_options[:extra_query]` is not supported for Realtime WebSocket " \ + "connections; omit it", + max_retries_message: "`request_options[:max_retries]` is not supported for Realtime WebSocket " \ + "connections; use 0 or omit it", + timeout_error: lambda do |url, cause| + OpenAI::Errors::RealtimeConnectionError.new(url: url, cause: cause) + end + ) + end + + private def validate_realtime_websocket_request!(websocket_base_url) if x509_identity?(@copy_options.fetch(:workload_identity)) raise OpenAI::Errors::Error, "X.509 workload identity does not support Realtime WebSocket connections" end @@ -184,106 +200,6 @@ def with_realtime_connection_request(path:, query:, websocket_base_url: nil, opt if websocket_base_url && @provider_runtime raise ArgumentError, "`websocket_base_url` cannot be combined with `provider`" end - - websocket_uri = parse_websocket_base_url(websocket_base_url) - - opts = options.to_h.dup - OpenAI::RequestOptions.validate!(opts) - extra_query = opts.delete(:extra_query) - unless extra_query.nil? || (extra_query.respond_to?(:empty?) && extra_query.empty?) - message = "`request_options[:extra_query]` is not supported for Realtime WebSocket " \ - "connections; omit it" - raise ArgumentError, message - end - - max_retries = opts[:max_retries] - unless max_retries.nil? || max_retries == 0 - message = "`request_options[:max_retries]` is not supported for Realtime WebSocket " \ - "connections; use 0 or omit it" - raise ArgumentError, message - end - - request = build_request( - { - method: :get, - path: path, - query: query, - security: {bearer_auth: true} - }, - opts - ) - error_request = if websocket_uri - with_websocket_base_url(request, path: path, base_url: websocket_uri) - else - request - end - - error_url = websocket_url(error_request.fetch(:url)) - - if @workload_identity_auth - deadline ||= request[:timeout]&.then do |timeout| - OpenAI::Internal::Util.monotonic_secs + timeout - end - end - - workload_identity_header = "Bearer #{OpenAI::Client::WORKLOAD_IDENTITY_API_KEY_PLACEHOLDER}" - if @workload_identity_auth && request.fetch(:headers)["authorization"] == workload_identity_header - token = @workload_identity_auth.get_token(deadline: deadline) - request = request.merge( - headers: request.fetch(:headers).merge("authorization" => "Bearer #{token}") - ) - end - - request = prepare_request(request, redirect_count: 0, retry_count: 0) - request = with_websocket_base_url(request, path: path, base_url: websocket_uri) if websocket_uri - - url = websocket_url(request.fetch(:url)) - headers = request.fetch(:headers).except("accept", "content-type").reject do |name, _value| - name.to_s.casecmp?("proxy-authorization") - end - - request = request.merge(url: url, headers: headers) - request = request_with_remaining_timeout(request, deadline) unless deadline.nil? - [request, deadline] - rescue Timeout::Error => e - raise( - OpenAI::Errors::RealtimeConnectionError.new( - url: error_url, - cause: e - ) - ) - end - - private def with_websocket_base_url(request, path:, base_url:) - url = OpenAI::Internal::Util.join_parsed_uri( - OpenAI::Internal::Util.parse_uri(base_url.to_s), - {path: OpenAI::Internal::Util.interpolate_path(path)} - ) - url.query = request.fetch(:url).query - request.merge(url: url) - end - - private def parse_websocket_base_url(value) - return if value.nil? - - uri = URI(value.to_s) - valid_scheme = %w[http https ws wss].include?(uri.scheme) - ambiguous_component = uri.userinfo || uri.query || uri.fragment - unless uri.absolute? && uri.host && valid_scheme && !ambiguous_component - message = "`websocket_base_url` must be an absolute HTTP or WebSocket URL " \ - "without credentials, query, or fragment" - raise ArgumentError, message - end - - uri - rescue URI::Error => e - raise ArgumentError, "`websocket_base_url` is not a valid URL", cause: e - end - - private def websocket_url(url) - url = url.dup - url.scheme = {"http" => "ws", "https" => "wss"}.fetch(url.scheme, url.scheme) - url end end end diff --git a/lib/openai/helpers/realtime/connection.rb b/lib/openai/helpers/realtime/connection.rb index 8a3df6568..496ebae5d 100644 --- a/lib/openai/helpers/realtime/connection.rb +++ b/lib/openai/helpers/realtime/connection.rb @@ -3,8 +3,8 @@ module OpenAI module Realtime # A live, typed Realtime WebSocket connection. - class Connection - include Enumerable + class Connection < OpenAI::WebSocket::Connection + include OpenAI::WebSocket::Protocol # @return [OpenAI::Realtime::ConnectionResources::Session] attr_reader :session @@ -20,8 +20,7 @@ class Connection # @api private def initialize(socket:, url:) - @socket = socket - @url = url + super @server_event_names = discriminator_values(OpenAI::Realtime::RealtimeServerEvent) @client_event_names = discriminator_values(OpenAI::Realtime::RealtimeClientEvent) @session = OpenAI::Realtime::ConnectionResources::Session.new(self) @@ -30,34 +29,6 @@ def initialize(socket:, url:) @input_audio_buffer = OpenAI::Realtime::ConnectionResources::InputAudioBuffer.new(self) end - # @return [URI::Generic] - attr_reader :url - - # Yield server events until the remote peer closes the connection. - def each - return enum_for(__method__) unless block_given? - - while (event = receive) - yield(event) - end - - self - end - - # Receive and parse the next server event, or return nil after a clean close. - def receive - data = receive_raw - return nil if data.nil? - - parse_event(data) - end - - # Receive the next raw WebSocket message. - def receive_raw - message = @socket.read - message&.to_str - end - # Parse raw JSON as a typed server event. Valid events that are newer than this # SDK remain observable as {UnknownServerEvent} values. def parse_event(data) @@ -115,65 +86,9 @@ def send_event(event) raise ArgumentError.new("Invalid Realtime client event."), cause: e end - # Send an already encoded text message. - def send_raw(data) - if closed? - raise( - OpenAI::Errors::RealtimeConnectionError.new( - url: @url, - message: "Cannot send on a closed Realtime WebSocket." - ) - ) - end - - text = data.dup - text.force_encoding(Encoding::UTF_8) if text.encoding == Encoding::BINARY - text = text.encode(Encoding::UTF_8) unless text.encoding == Encoding::UTF_8 - unless text.valid_encoding? - raise ArgumentError, "Realtime WebSocket text must contain valid UTF-8" - end - - @socket.write(text) - nil - end - - # Close the connection. - def close(code: 1000, reason: "") - return if closed? - - @socket.close(code: code, reason: reason) - nil - end - - # Abort without waiting for the WebSocket close handshake. - # - # @api private - def abort - return if closed? - - @socket.abort - nil - end - - # @return [Boolean] - def closed? = @socket.closed? - - private def discriminator_values(union) - union.variants.to_h do |variant| - value = variant.fields.fetch(:type).fetch(:const) - [value.to_s, true] - end - end - private def event_type(event) - unless event.is_a?(Hash) - raise ArgumentError, "Realtime server event must be a JSON object" - end - - type = event[:type] - return type if type.is_a?(String) || type.is_a?(Symbol) - - raise ArgumentError, "Realtime server event type must be a string or symbol" + super(event, message: "Realtime server event must be a JSON object") unless event.is_a?(Hash) + super(event, message: "Realtime server event type must be a string or symbol") end private def validate_discriminator!(event, allowed, kind:) @@ -194,6 +109,14 @@ def closed? = @socket.closed? ArgumentError.new("Realtime event is missing required fields or contains invalid values") end + + private def connection_error(message) + OpenAI::Errors::RealtimeConnectionError.new(url: @url, message: message) + end + + private def closed_send_message = "Cannot send on a closed Realtime WebSocket." + + private def invalid_text_message = "Realtime WebSocket text must contain valid UTF-8" end end end diff --git a/lib/openai/helpers/realtime/connection_manager.rb b/lib/openai/helpers/realtime/connection_manager.rb index 510b71d5c..4a020f257 100644 --- a/lib/openai/helpers/realtime/connection_manager.rb +++ b/lib/openai/helpers/realtime/connection_manager.rb @@ -5,19 +5,7 @@ module Realtime # Internal block-scoped lifecycle manager for Realtime WebSocket connections. # # @api private - class ConnectionManager - RESERVED_TRANSPORT_OPTIONS = [ - :alpn_protocols, - :headers, - :hostname, - :port, - :protocol, - :scheme, - :ssl_context, - :timeout, - :url - ].freeze - + class ConnectionManager < OpenAI::WebSocket::ConnectionManager # @api private def initialize( client:, @@ -28,79 +16,36 @@ def initialize( transport_options:, connection_class: OpenAI::Realtime::Connection ) - @client = client - @query = query + query = query .to_h .to_h do |key, value| [key.to_s.dup.freeze, value.to_s.dup.freeze] end .freeze - @websocket_base_url = websocket_base_url&.to_s&.dup&.freeze - @transport = transport - @request_options = request_options - @connection_class = connection_class - transport_options = transport_options.dup.freeze - reserved_options = transport_options.keys.select do |key| - (key.is_a?(String) || key.is_a?(Symbol)) && RESERVED_TRANSPORT_OPTIONS.include?(key.to_sym) - end - - unless reserved_options.empty? - raise( - ArgumentError, - "`transport_options` cannot include #{reserved_options.map(&:inspect).join(", ")}" + base_url = websocket_base_url&.to_s&.dup&.freeze + request = lambda do |&request_block| + client.with_realtime_connection_request( + path: "realtime", + query: query, + websocket_base_url: base_url, + options: request_options, + &request_block ) end - @transport_options = transport_options - end - - # Open the WebSocket and yield a typed connection for the lifetime of the block. - # - # @api private - # - # @yieldparam connection [OpenAI::Realtime::Connection] - # @return [Object] - def open - raise ArgumentError, "A block is required to open a Realtime WebSocket." unless block_given? - - transport = @transport || OpenAI::Realtime::Transports::AsyncWebSocket.new - unless transport.respond_to?(:open) - raise ArgumentError, "`transport` must respond to `open`" - end - - @client - .with_realtime_connection_request( - path: "realtime", - query: @query, - websocket_base_url: @websocket_base_url, - options: @request_options - ) do |request, mark_handshake_completed| - transport - .open( - url: request.fetch(:url), - headers: request.fetch(:headers), - timeout: request.fetch(:timeout), - **@transport_options - ) do |socket| - mark_handshake_completed.call - connection = @connection_class.new(socket: socket, url: request.fetch(:url)) - begin - yield(connection) - ensure - pending_error = $ERROR_INFO - begin - if pending_error - connection.abort unless connection.closed? - else - connection.close unless connection.closed? - end - - rescue StandardError - raise if pending_error.nil? - end - end - end + super( + transport: transport, + transport_options: transport_options, + default_transport: -> { OpenAI::Realtime::Transports::AsyncWebSocket.new }, + connection_class: connection_class, + request: request, + block_error_message: "A block is required to open a Realtime WebSocket.", + abort_after_block: -> (_connection, pending_error) { !pending_error.nil? }, + transport_error_message: "`transport` must respond to `open`", + reserved_options_error_message: lambda do |reserved| + "`transport_options` cannot include #{reserved.map(&:inspect).join(", ")}" end + ) end end end diff --git a/lib/openai/helpers/realtime/errors.rb b/lib/openai/helpers/realtime/errors.rb index d74576e9a..2e7d8767e 100644 --- a/lib/openai/helpers/realtime/errors.rb +++ b/lib/openai/helpers/realtime/errors.rb @@ -3,32 +3,8 @@ module OpenAI module Errors # Raised when a Realtime WebSocket cannot be opened or used. - class RealtimeConnectionError < OpenAI::Errors::Error - # @return [URI::Generic] - attr_reader :url - - # HTTP status returned by a failed WebSocket upgrade, when available. - # - # @api private - # - # @return [Integer, nil] - attr_reader :http_status - - # @return [Exception, nil] - def cause = @cause.nil? ? super : @cause - - # @api private - # - # @param url [URI::Generic] - # @param message [String, nil] - # @param cause [Exception, nil] - # @param http_status [Integer, nil] - def initialize(url:, message: nil, cause: nil, http_status: nil) - @url = sanitized_error_url(url) - @cause = cause - @http_status = http_status - super(message || "Realtime WebSocket connection error.") - end + class RealtimeConnectionError < OpenAI::Errors::WebSocketConnectionError + private def default_message = "Realtime WebSocket connection error." private def sanitized_error_url(url) query = url.query @@ -52,7 +28,7 @@ def initialize(url:, message: nil, cause: nil, http_status: nil) end # Raised when a Realtime WebSocket message cannot be parsed as a typed event. - class RealtimeProtocolError < OpenAI::Errors::Error + class RealtimeProtocolError < OpenAI::Errors::WebSocketProtocolError # @return [String] attr_reader :data diff --git a/lib/openai/helpers/realtime/transports/async_websocket.rb b/lib/openai/helpers/realtime/transports/async_websocket.rb index 7f343130f..424ddbd06 100644 --- a/lib/openai/helpers/realtime/transports/async_websocket.rb +++ b/lib/openai/helpers/realtime/transports/async_websocket.rb @@ -3,384 +3,41 @@ module OpenAI module Realtime module Transports - # Default Realtime transport backed by the optional async-websocket gem. - class AsyncWebSocket - REDACTED_HEADER_VALUE = "[REDACTED]" - private_constant :REDACTED_HEADER_VALUE - - # Preserve real header fields for the wire while presenting a safe snapshot to - # Protocol::HTTP1 tracing, which calls #to_h before writing the request. - class TraceSafeHeaderFields - include Enumerable - - def initialize(fields) - @fields = fields.to_a.freeze - end - - def each(&block) = @fields.each(&block) - - def to_h - @fields.to_h do |name, value| - safe_value = if OpenAI::Internal::Logging.sensitive_header?(name) - REDACTED_HEADER_VALUE - else - value - end - - [name, safe_value] - end - end - end - - private_constant :TraceSafeHeaderFields - - module TraceSafeHeaders - def header(&block) - return super(&block) if block - - TraceSafeHeaderFields.new(super) - end - end - - private_constant :TraceSafeHeaders - - # HTTP/1 tracing observes the request target passed to #write_request. - # Keep that target redacted while restoring the real call ID only at the - # private connection's first serialized request-line boundary. - class TraceSafeRequestStream - def initialize(stream, trace_request_line:, wire_request_line:) - @stream = stream - @trace_request_line = trace_request_line - @wire_request_line = wire_request_line - @first_write = true - end - - def write(data) - if @first_write - @first_write = false - unless data == @trace_request_line - raise IOError, "Unexpected Realtime WebSocket request-line serialization." - end - - data = @wire_request_line - end - - @stream.write(data) - end - end - - private_constant :TraceSafeRequestStream - - module TraceSafeRequestWriter - def write_request(authority, method, target, version, headers) - trace_target = target.gsub(/([?&])([^&=]+)(?:=[^&]*)?/) do |parameter| - separator = Regexp.last_match(1) - name = Regexp.last_match(2) - if URI.decode_www_form_component(name) == "call_id" - "#{separator}#{name}=[REDACTED]" - else - parameter - end - end - - original_stream = @stream - @stream = TraceSafeRequestStream.new( - original_stream, - trace_request_line: "#{method} #{trace_target} #{version}\r\n", - wire_request_line: "#{method} #{target} #{version}\r\n" - ) - - super(authority, method, trace_target, version, headers) - ensure - @stream = original_stream - end - end - - private_constant :TraceSafeRequestWriter - - class TraceSafeProtocol - def initialize(protocol) - @protocol = protocol - end - - def client(peer, **options) - @protocol.client(peer, **options).extend(TraceSafeRequestWriter) - end - - def to_s = @protocol.to_s - end - - private_constant :TraceSafeProtocol - - # Async's ordinary framer close flushes buffered output. Exceptional cleanup - # must instead close the raw socket first, then release the acquired pool slot. - module AbortableFramer - def abort - stream = @stream - pool = @pool - connection = @connection - @pool = nil - @connection = nil - - begin - io = stream.to_io - io.close unless io.closed? - ensure - pool&.release(connection) - end - end + # Compatibility facade for the shared async WebSocket transport. + class AsyncWebSocket < OpenAI::WebSocket::AsyncWebSocketTransport + ERROR_FACTORY = lambda do |url:, message: nil, cause: nil, http_status: nil| + OpenAI::Errors::RealtimeConnectionError.new( + url: url, + message: message, + cause: cause, + http_status: http_status + ) end - private_constant :AbortableFramer - - # Configure the native TLS context used by secure Realtime WebSockets. - # Peer and hostname verification and HTTP/1.1 ALPN remain SDK-owned. - def initialize(&tls_configurator) - @tls_configurator = tls_configurator - end + private_constant :ERROR_FACTORY - class Socket - # @api private + # Compatibility wrapper for the socket class yielded before the shared + # transport extraction. + # + # @api private + class Socket < OpenAI::WebSocket::AsyncWebSocketTransport::Socket def initialize(connection, url:) - @connection = connection - @url = url - @aborted = false - end - - def read - @connection.read - rescue StandardError => e - raise OpenAI::Errors::RealtimeConnectionError.new(url: @url, cause: e) + super(connection, url: url, error_factory: ERROR_FACTORY) end - - def write(message) - @connection.write(message) - @connection.flush - rescue StandardError => e - raise OpenAI::Errors::RealtimeConnectionError.new(url: @url, cause: e) - end - - def close(code: 1000, reason: "") - @connection.close(code, reason) - rescue StandardError => e - raise OpenAI::Errors::RealtimeConnectionError.new(url: @url, cause: e) - end - - # @api private - def abort - framer = @connection.framer - framer.extend(AbortableFramer) - framer.abort - @aborted = true - rescue StandardError => e - raise OpenAI::Errors::RealtimeConnectionError.new(url: @url, cause: e) - end - - # @api private - def aborted? = @aborted - - def closed? = @aborted || @connection.closed? end - def open(url:, headers:, timeout:, **endpoint_options) - load_dependencies(url) - - # Proxy credentials belong only on the CONNECT request assembled from - # proxy configuration. Never forward a caller-supplied value to the - # Realtime origin handshake. - headers = headers.reject do |name, _value| - name.to_s.casecmp?("proxy-authorization") - end - - # Classic WebSocket negotiation uses HTTP/1.1. Pinning ALPN also avoids an - # HTTP/2 selection on servers that advertise both protocols. - options = { - alpn_protocols: ::Async::HTTP::Protocol::HTTP11.names, - **endpoint_options - } - if url.scheme == "wss" || @tls_configurator - options[:ssl_context] = build_tls_context(url) - end - - sideband = sideband_call?(url) - endpoint_url = url.dup - endpoint_url.query = nil if sideband - - endpoint = ::Async::HTTP::Endpoint.parse( - endpoint_url.to_s, - **options - ) - request_target = endpoint.path - request_target = "#{request_target}?#{url.query}" if sideband - proxy_client = nil - if (proxy = proxy_uri(url)) - proxy_endpoint = ::Async::HTTP::Endpoint.parse(proxy_url(proxy).to_s) - proxy_client = ::Async::HTTP::Client.open(proxy_endpoint) - tunnel = ::Async::HTTP::Proxy.new( - proxy_client, - authority(url, include_default_port: true), - trace_safe_headers(proxy_headers(proxy)) - ) - endpoint = tunnel.wrap_endpoint(endpoint) - end - - block_error = nil - # Keep the request timeout scoped to WebSocket negotiation. Endpoint timeouts - # remain installed on the socket and would otherwise terminate healthy idle - # Realtime sessions after the ordinary HTTP request timeout. - ::Kernel.Sync() do - client_options = sideband ? {protocol: TraceSafeProtocol.new(endpoint.protocol)} : {} - client = ::Async::WebSocket::Client.open(endpoint, **client_options) - connection = nil - socket = nil - begin - connection = negotiate( - client, - endpoint, - request_target: request_target, - headers: headers, - timeout: timeout - ) - socket = Socket.new(connection, url: url) - begin - yield(socket) - rescue StandardError => e - block_error = e - raise - end - - ensure - close_resources( - connection, - client, - proxy_client, - connection_aborted: socket&.aborted?, - pending_error: $ERROR_INFO - ) - end - end - - rescue OpenAI::Errors::RealtimeConnectionError - raise - rescue StandardError => e - raise if e.equal?(block_error) - - raise( - OpenAI::Errors::RealtimeConnectionError.new( - url: url, - cause: e, - http_status: handshake_status(e) - ) + def initialize(&tls_configurator) + super( + product_name: "Realtime", + error_class: OpenAI::Errors::RealtimeConnectionError, + error_factory: ERROR_FACTORY, + sensitive_query_parameter: "call_id", +&tls_configurator ) end - private def negotiate(client, endpoint, request_target:, headers:, timeout:) - safe_headers = trace_safe_headers(headers) - operation = lambda do - client.connect(authority(endpoint.url), request_target, headers: safe_headers) - end - - return operation.call if timeout.nil? - - ::Async::Task.current.with_timeout(timeout, &operation) - end - - private def sideband_call?(url) - return false unless url.query - - URI.decode_www_form(url.query).any? { |name, _value| name == "call_id" } - end - - private def build_tls_context(url) - unless url.scheme == "wss" - raise ArgumentError, "TLS configuration requires a wss:// Realtime endpoint" - end - - context = OpenSSL::SSL::SSLContext.new - @tls_configurator&.call(context) - if context.verify_callback - raise ArgumentError, "Realtime WebSocket TLS configuration cannot set verify_callback" - end - - context.set_params(verify_mode: OpenSSL::SSL::VERIFY_PEER) - context.verify_hostname = true - context.alpn_protocols = ::Async::HTTP::Protocol::HTTP11.names - context - end - - private def close_resources(connection, client, proxy_client, connection_aborted:, pending_error:) - cleanup_error = nil - resources = [client, proxy_client] - resources.unshift(connection) unless connection_aborted - resources.compact.each do |resource| - resource.close - rescue StandardError => e - cleanup_error ||= e - end - - raise cleanup_error if pending_error.nil? && cleanup_error - end - - private def load_dependencies(url) - require("async/websocket/client") - require("async/http/endpoint") - require("async/http/proxy") - - rescue LoadError => e - message = "Realtime WebSockets require the `async-websocket` gem. " \ - "Add `gem \"async-websocket\"` to your Gemfile." - raise OpenAI::Errors::RealtimeConnectionError.new(url: url, message: message, cause: e) - end - - private def authority(url, include_default_port: false) - host = url.hostname - host = "[#{host}]" if host.include?(":") - default_port = %w[https wss].include?(url.scheme) ? 443 : 80 - return host if !include_default_port && url.port == default_port - - "#{host}:#{url.port}" - end - - private def proxy_uri(url) - policy_url = url.dup - policy_url.scheme = {"ws" => "http", "wss" => "https"}.fetch(url.scheme, url.scheme) - policy_url.find_proxy - end - - private def proxy_url(proxy) - unless %w[http https].include?(proxy.scheme) && proxy.hostname - raise ArgumentError, "Realtime WebSocket proxy must be an absolute HTTP or HTTPS URL" - end - - proxy.dup.tap do |url| - url.user = nil - url.password = nil - end - end - - private def proxy_headers(proxy) - return {} unless proxy.user - - user = URI::RFC2396_PARSER.unescape(proxy.user) - password = URI::RFC2396_PARSER.unescape(proxy.password.to_s) - {"proxy-authorization" => "Basic #{["#{user}:#{password}"].pack("m0")}"} - end - - private def trace_safe_headers(headers) - fields = ::Protocol::HTTP::Headers[headers].to_a - trace_safe_headers_class.new(fields) - end - - private def trace_safe_headers_class - @trace_safe_headers_class ||= Class.new(::Protocol::HTTP::Headers) do - include(TraceSafeHeaders) - end - end - - private def handshake_status(error) - return unless error.is_a?(::Async::WebSocket::ConnectionError) - - error.response.status + private def build_socket(connection, url:) + Socket.new(connection, url: url) end end end diff --git a/lib/openai/helpers/responses_websocket.rb b/lib/openai/helpers/responses_websocket.rb new file mode 100644 index 000000000..a7698f5e0 --- /dev/null +++ b/lib/openai/helpers/responses_websocket.rb @@ -0,0 +1,12 @@ +# frozen_string_literal: true + +# Responses WebSocket mode is a custom SDK runtime layered on generated +# Responses protocol models. Generated HTTP resources remain generator-owned. +require_relative "responses_websocket/errors" +require_relative "responses_websocket/unknown_server_event" +require_relative "responses_websocket/connection_resources" +require_relative "responses_websocket/connection" +require_relative "responses_websocket/connection_manager" +require_relative "responses_websocket/transports/async_websocket" +require_relative "responses_websocket/client_extension" +require_relative "responses_websocket/resources/responses_extension" diff --git a/lib/openai/helpers/responses_websocket/client_extension.rb b/lib/openai/helpers/responses_websocket/client_extension.rb new file mode 100644 index 000000000..d284d5298 --- /dev/null +++ b/lib/openai/helpers/responses_websocket/client_extension.rb @@ -0,0 +1,62 @@ +# frozen_string_literal: true + +module OpenAI + module Helpers + module ResponsesWebSocket + # Responses WebSocket request integration kept outside generated client code. + module ClientExtension + include OpenAI::WebSocket::ClientRequest + + # @api private + def with_responses_websocket_connection_request(websocket_base_url: nil, options: nil) + build = lambda do |deadline| + build_responses_websocket_connection_request( + websocket_base_url: websocket_base_url, + options: options, + deadline: deadline + ) + end + + with_websocket_connection_retry( + error_class: OpenAI::Errors::ResponsesConnectionError, + build: build + ) { |request, marker| yield(request, marker) } + end + + private def build_responses_websocket_connection_request( + websocket_base_url:, + options:, + deadline: nil + ) + build_shared_websocket_connection_request( + path: "responses", + query: {}, + websocket_base_url: websocket_base_url, + options: options, + deadline: deadline, + validate: -> (_value) { validate_responses_websocket_request! }, + invalid_base_url_message: "websocket_base_url must be an absolute HTTP or WebSocket URL " \ + "without credentials, query, or fragment", + malformed_base_url_message: "websocket_base_url is not a valid URL", + preserve_base_url_cause: false, + extra_query_message: "request_options extra_query is not supported for Responses WebSocket connections", + max_retries_message: "request_options max_retries is not supported for Responses WebSocket connections", + timeout_error: -> (url, _cause) { OpenAI::Errors::ResponsesConnectionError.new(url: url) } + ) + end + + private def validate_responses_websocket_request! + if x509_identity?(@copy_options.fetch(:workload_identity)) + raise OpenAI::Errors::Error, "X.509 workload identity does not support Responses WebSocket connections" + end + + if @provider_runtime + raise OpenAI::Errors::Error, "Responses WebSocket connections are not supported by providers." + end + end + end + end + end +end + +OpenAI::Client.include(OpenAI::Helpers::ResponsesWebSocket::ClientExtension) diff --git a/lib/openai/helpers/responses_websocket/connection.rb b/lib/openai/helpers/responses_websocket/connection.rb new file mode 100644 index 000000000..257ed5d4e --- /dev/null +++ b/lib/openai/helpers/responses_websocket/connection.rb @@ -0,0 +1,254 @@ +# frozen_string_literal: true + +module OpenAI + module Responses + # A live, typed Responses WebSocket connection. + class Connection < OpenAI::WebSocket::Connection + include OpenAI::WebSocket::Protocol + + # @return [OpenAI::Responses::ConnectionResources::Response] + attr_reader :response + + # @api private + def initialize(socket:, url:) + super + @poisoned = false + @state = :open + @owner_thread = Thread.current.object_id + @reading = false + @server_event_names = discriminator_values(OpenAI::Responses::ResponsesServerEvent) + @response = OpenAI::Responses::ConnectionResources::Response.new(self) + end + + def each + return enum_for(__method__) unless block_given? + + with_read_lease do + while (event = read_one) + yield(event) + end + end + + self + end + + def receive + with_read_lease { read_one } + end + + # Validate, encode, and send a Responses client event. + # + # @return [nil] + def send_event(event) + assert_owner! + raise connection_error("Cannot send on a closed Responses WebSocket.") if closed? || poisoned? + + write_encoded(JSON.generate(encode_client_event(event))) + rescue OpenAI::Errors::ResponsesConnectionError, + OpenAI::Errors::ResponsesClientEventError, + OpenAI::Errors::ResponsesSendError + raise + rescue StandardError + raise OpenAI::Errors::ResponsesClientEventError.new, cause: nil + end + + # @api private + def poisoned? = @poisoned + + def close(code: 1000, reason: "") + assert_owner! + return if closed? + return abort if poisoned? + + super + @state = :closed + nil + rescue OpenAI::Errors::ResponsesConnectionError + raise + rescue StandardError + @state = :closed + raise OpenAI::Errors::ResponsesConnectionError.new(url: @url), cause: nil + end + + # @api private + def abort + assert_owner! + return if closed? + + super + @state = :closed + nil + rescue OpenAI::Errors::ResponsesConnectionError + raise + rescue StandardError + @state = :closed + raise OpenAI::Errors::ResponsesConnectionError.new(url: @url), cause: nil + end + + def closed? = @state == :closed || super + + private def parse_event(data) + parsed = JSON.parse(data, symbolize_names: true) + type = event_type(parsed, message: "Responses server event must be a JSON object") + unless @server_event_names.key?(type.to_s) + return OpenAI::Responses::UnknownServerEvent.new(data: parsed) + end + + state = OpenAI::Internal::Type::Converter.new_coerce_state + OpenAI::Internal::Type::Converter.coerce( + OpenAI::Responses::ResponsesServerEvent, + parsed, + state: state + ) + rescue OpenAI::Errors::ResponsesProtocolError + raise + rescue StandardError + raise OpenAI::Errors::ResponsesProtocolError.new, cause: nil + end + + private def read_one + assert_owner! + return nil if closed? + + data = receive_raw + if data.nil? + @state = :closed + return nil + end + + parse_event(data.to_str) + rescue OpenAI::Errors::ResponsesProtocolError, OpenAI::Errors::ResponsesConnectionError + raise + rescue StandardError + raise OpenAI::Errors::ResponsesConnectionError.new(url: @url), cause: nil + end + + private def encode_client_event(event) + validate_event_tree!(event) + payload = if event.is_a?(Hash) + OpenAI::Internal::Type::Unknown.dump(event, state: {can_retry: true}) + else + OpenAI::Internal::Type::Converter.dump(OpenAI::Responses::ResponsesClientEvent, event) + end + + raise OpenAI::Errors::ResponsesClientEventError.new unless payload.is_a?(Hash) + + normalize_event_keys(payload) + rescue OpenAI::Errors::ResponsesClientEventError + raise + rescue StandardError + raise OpenAI::Errors::ResponsesClientEventError.new, cause: nil + end + + private def write_encoded(data) + send_raw(data) + nil + rescue StandardError + @poisoned = true + raise OpenAI::Errors::ResponsesSendError.new(url: @url), cause: nil + end + + private def with_read_lease + acquired = false + assert_owner! + if @reading + raise connection_error("Responses WebSocket already has an active reader.") + end + + @reading = true + acquired = true + yield + ensure + @reading = false if acquired + end + + private def assert_owner! + return if @owner_thread == Thread.current.object_id + + raise connection_error("Responses WebSocket connections are single-owner.") + end + + private def validate_event_tree!(value, active = {}.compare_by_identity) + case value + when OpenAI::Internal::Type::BaseModel + guard_cycle!(value, active) + begin + data = value.to_h + reject_duplicate_serialized_model_keys!(value.class, data) + validate_event_tree!(data, active) + ensure + active.delete(value) + end + + when Hash + guard_cycle!(value, active) + reject_duplicate_semantic_keys!(value) + begin + value.each_value { |item| validate_event_tree!(item, active) } + ensure + active.delete(value) + end + + when Array + guard_cycle!(value, active) + begin + value.each { |item| validate_event_tree!(item, active) } + ensure + active.delete(value) + end + end + + nil + end + + private def reject_duplicate_serialized_model_keys!(model, data) + serialized_keys = data.keys.map do |key| + name = key.is_a?(String) ? key.to_sym : key + field = model.known_fields[name] + field ? field.fetch(:api_name) : name + end + + return if serialized_keys.map(&:to_s).uniq.length == serialized_keys.length + + raise OpenAI::Errors::ResponsesClientEventError.new + end + + private def guard_cycle!(value, active) + raise OpenAI::Errors::ResponsesClientEventError.new if active.key?(value) + + active[value] = true + end + + private def reject_duplicate_semantic_keys!(payload) + keys = payload.keys + unless keys.all? { |key| key.is_a?(String) || key.is_a?(Symbol) } + raise OpenAI::Errors::ResponsesClientEventError.new + end + + return if keys.map(&:to_s).uniq.length == keys.length + + raise OpenAI::Errors::ResponsesClientEventError.new + end + + private def normalize_event_keys(value) + case value + when Hash + value.to_h do |key, item| + normalized_key = key.is_a?(String) ? key.to_sym : key + [normalized_key, normalize_event_keys(item)] + end + + when Array + value.map { |item| normalize_event_keys(item) } + else + value + end + end + + private def connection_error(message) + OpenAI::Errors::ResponsesConnectionError.new(url: @url, message: message) + end + + end + end +end diff --git a/lib/openai/helpers/responses_websocket/connection_manager.rb b/lib/openai/helpers/responses_websocket/connection_manager.rb new file mode 100644 index 000000000..502d1acb0 --- /dev/null +++ b/lib/openai/helpers/responses_websocket/connection_manager.rb @@ -0,0 +1,34 @@ +# frozen_string_literal: true + +module OpenAI + module Responses + # Internal block-scoped lifecycle manager for Responses WebSocket connections. + # + # @api private + class ConnectionManager < OpenAI::WebSocket::ConnectionManager + # @api private + def initialize(client:, websocket_base_url:, transport:, request_options:, transport_options:) + base_url = websocket_base_url&.to_s&.dup&.freeze + request = lambda do |&request_block| + client.with_responses_websocket_connection_request( + websocket_base_url: base_url, + options: request_options, + &request_block + ) + end + + super( + transport: transport, + transport_options: transport_options, + default_transport: -> { OpenAI::Responses::Transports::AsyncWebSocket.new }, + connection_class: OpenAI::Responses::Connection, + request: request, + block_error_message: "A block is required to open a Responses WebSocket.", + abort_after_block: lambda do |connection, pending_error| + !pending_error.nil? || connection.poisoned? + end + ) + end + end + end +end diff --git a/lib/openai/helpers/responses_websocket/connection_resources.rb b/lib/openai/helpers/responses_websocket/connection_resources.rb new file mode 100644 index 000000000..d187157fa --- /dev/null +++ b/lib/openai/helpers/responses_websocket/connection_resources.rb @@ -0,0 +1,22 @@ +# frozen_string_literal: true + +module OpenAI + module Responses + module ConnectionResources + class Response + # @api private + def initialize(connection) + @connection = connection + end + + # Send a response.create event over this connection. + # + # @return [nil] + def create(**params) + event_params = params.reject { |key, _value| key == :type || key == "type" } + @connection.send_event(**event_params, type: :"response.create") + end + end + end + end +end diff --git a/lib/openai/helpers/responses_websocket/errors.rb b/lib/openai/helpers/responses_websocket/errors.rb new file mode 100644 index 000000000..527079bec --- /dev/null +++ b/lib/openai/helpers/responses_websocket/errors.rb @@ -0,0 +1,56 @@ +# frozen_string_literal: true + +module OpenAI + module Errors + # Raised when a Responses WebSocket cannot be opened or used. + class ResponsesConnectionError < OpenAI::Errors::WebSocketConnectionError + # @return [Exception, nil] + def cause = nil + + private def default_message = "Responses WebSocket connection error." + + private def sanitized_error_url(url) + sanitized = url.dup + sanitized.user = nil if sanitized.respond_to?(:user=) + sanitized.password = nil if sanitized.respond_to?(:password=) + sanitized.query = nil if sanitized.respond_to?(:query=) + sanitized.fragment = nil if sanitized.respond_to?(:fragment=) + sanitized + rescue ArgumentError, URI::Error + URI("wss://invalid") + end + end + + # Raised when a Responses WebSocket message is malformed. + class ResponsesProtocolError < OpenAI::Errors::WebSocketProtocolError + def cause = nil + + # @api private + def initialize + super("Invalid Responses WebSocket event.") + end + end + + # Raised before a malformed client event can be written. + class ResponsesClientEventError < OpenAI::Errors::Error + def cause = nil + + # @api private + def initialize + super("Invalid Responses WebSocket client event.") + end + end + + # Raised when a write may or may not have reached the server. + class ResponsesSendError < ResponsesConnectionError + # @return [Symbol, :unknown] + attr_reader :outcome + + # @api private + def initialize(url:) + @outcome = :unknown + super(url: url, message: "Responses WebSocket send outcome is unknown.") + end + end + end +end diff --git a/lib/openai/helpers/responses_websocket/resources/responses_extension.rb b/lib/openai/helpers/responses_websocket/resources/responses_extension.rb new file mode 100644 index 000000000..c9a645060 --- /dev/null +++ b/lib/openai/helpers/responses_websocket/resources/responses_extension.rb @@ -0,0 +1,31 @@ +# frozen_string_literal: true + +module OpenAI + module Helpers + module ResponsesWebSocket + # Ruby-native WebSocket entry point layered onto the generated Responses resource. + module Connections + def connect( + websocket_base_url: nil, + request_options: nil, + transport: nil, + transport_options: {}, + &block + ) + raise ArgumentError, "A block is required to open a Responses WebSocket." unless block + + manager = OpenAI::Responses::ConnectionManager.new( + client: @client, + websocket_base_url: websocket_base_url, + transport: transport, + request_options: request_options, + transport_options: transport_options + ) + manager.open(&block) + end + end + end + end +end + +OpenAI::Resources::Responses.include(OpenAI::Helpers::ResponsesWebSocket::Connections) diff --git a/lib/openai/helpers/responses_websocket/transports/async_websocket.rb b/lib/openai/helpers/responses_websocket/transports/async_websocket.rb new file mode 100644 index 000000000..b8b1dd943 --- /dev/null +++ b/lib/openai/helpers/responses_websocket/transports/async_websocket.rb @@ -0,0 +1,29 @@ +# frozen_string_literal: true + +module OpenAI + module Responses + module Transports + # Responses facade for the shared optional async WebSocket transport. + # + # @api private + class AsyncWebSocket < OpenAI::WebSocket::AsyncWebSocketTransport + def initialize + error_factory = lambda do |url:, message: nil, http_status: nil, **_options| + OpenAI::Errors::ResponsesConnectionError.new( + url: url, + message: message, + http_status: http_status + ) + end + + super( + product_name: "Responses", + error_class: OpenAI::Errors::ResponsesConnectionError, + error_factory: error_factory, + dependency_message: "Responses WebSockets require the async-websocket gem. Add it to your Gemfile." + ) + end + end + end + end +end diff --git a/lib/openai/helpers/responses_websocket/unknown_server_event.rb b/lib/openai/helpers/responses_websocket/unknown_server_event.rb new file mode 100644 index 000000000..e951f7648 --- /dev/null +++ b/lib/openai/helpers/responses_websocket/unknown_server_event.rb @@ -0,0 +1,54 @@ +# frozen_string_literal: true + +module OpenAI + module Responses + # A valid JSON event whose discriminator is newer than this SDK version. + class UnknownServerEvent + # @return [Symbol] + attr_reader :type + + # @return [Hash{Symbol=>Object}] + attr_reader :data + + # @return [Object] + attr_reader :stream_id + + # @api private + def initialize(data:) + value = data.fetch(:type) + unless value.is_a?(String) || value.is_a?(Symbol) + raise ArgumentError, "Responses server event type must be a string or symbol" + end + + @type = value.to_sym + @stream_id = data[:stream_id] + @data = freeze_json(data) + freeze + end + + # @return [Hash{Symbol=>Object}] + def to_h = @data + + # Keep routine diagnostics payload-free. Callers that explicitly need the + # unknown JSON can use {#data} or {#to_h}. + def inspect = "#<#{self.class} type=#{@type.inspect}>" + + alias to_s inspect + + private def freeze_json(value) + case value + when Hash + value.each do |key, item| + freeze_json(key) + freeze_json(item) + end + + when Array + value.each { |item| freeze_json(item) } + end + + value.freeze + end + end + end +end diff --git a/lib/openai/helpers/websocket.rb b/lib/openai/helpers/websocket.rb new file mode 100644 index 000000000..886c44112 --- /dev/null +++ b/lib/openai/helpers/websocket.rb @@ -0,0 +1,10 @@ +# frozen_string_literal: true + +# Internal, product-neutral WebSocket runtime shared by SDK WebSocket APIs. +# Product helpers keep their event models and resource facades separate. +require_relative "websocket/errors" +require_relative "websocket/connection" +require_relative "websocket/connection_manager" +require_relative "websocket/protocol" +require_relative "websocket/async_websocket_transport" +require_relative "websocket/client_request" diff --git a/lib/openai/helpers/websocket/async_websocket_transport.rb b/lib/openai/helpers/websocket/async_websocket_transport.rb new file mode 100644 index 000000000..09e394445 --- /dev/null +++ b/lib/openai/helpers/websocket/async_websocket_transport.rb @@ -0,0 +1,415 @@ +# frozen_string_literal: true + +module OpenAI + module WebSocket + # Product-neutral transport backed by the optional async-websocket gem. + # + # @api private + class AsyncWebSocketTransport + REDACTED_HEADER_VALUE = "[REDACTED]" + private_constant :REDACTED_HEADER_VALUE + + # Preserve real header fields for the wire while presenting a safe snapshot to + # Protocol::HTTP1 tracing, which calls #to_h before writing the request. + class TraceSafeHeaderFields + include Enumerable + + def initialize(fields) + @fields = fields.to_a.freeze + end + + def each(&block) = @fields.each(&block) + + def to_h + @fields.to_h do |name, value| + safe_value = if OpenAI::Internal::Logging.sensitive_header?(name) + REDACTED_HEADER_VALUE + else + value + end + + [name, safe_value] + end + end + end + + private_constant :TraceSafeHeaderFields + + module TraceSafeHeaders + def header(&block) + return super(&block) if block + + TraceSafeHeaderFields.new(super) + end + end + + private_constant :TraceSafeHeaders + + # HTTP/1 tracing observes the request target passed to #write_request. + # Keep that target redacted while restoring the real call ID only at the + # private connection's first serialized request-line boundary. + class TraceSafeRequestStream + def initialize(stream, trace_request_line:, wire_request_line:) + @stream = stream + @trace_request_line = trace_request_line + @wire_request_line = wire_request_line + @first_write = true + end + + def write(data) + if @first_write + @first_write = false + unless data == @trace_request_line + raise IOError, "Unexpected WebSocket request-line serialization." + end + + data = @wire_request_line + end + + @stream.write(data) + end + end + + private_constant :TraceSafeRequestStream + + module TraceSafeRequestWriter + def write_request(authority, method, target, version, headers) + return super unless @sensitive_query_parameter + + trace_target = target.gsub(/([?&])([^&=]+)(?:=[^&]*)?/) do |parameter| + separator = Regexp.last_match(1) + name = Regexp.last_match(2) + if URI.decode_www_form_component(name) == @sensitive_query_parameter + "#{separator}#{name}=[REDACTED]" + else + parameter + end + end + + original_stream = @stream + @stream = TraceSafeRequestStream.new( + original_stream, + trace_request_line: "#{method} #{trace_target} #{version}\r\n", + wire_request_line: "#{method} #{target} #{version}\r\n" + ) + + super(authority, method, trace_target, version, headers) + ensure + @stream = original_stream + end + end + + private_constant :TraceSafeRequestWriter + + class TraceSafeProtocol + def initialize(protocol, sensitive_query_parameter:) + @protocol = protocol + @sensitive_query_parameter = sensitive_query_parameter + end + + def client(peer, **options) + @protocol.client(peer, **options).tap do |client| + client.extend(TraceSafeRequestWriter) + client.instance_variable_set(:@sensitive_query_parameter, @sensitive_query_parameter) + end + end + + def to_s = @protocol.to_s + end + + private_constant :TraceSafeProtocol + + # Async's ordinary framer close flushes buffered output. Exceptional cleanup + # must instead close the raw socket first, then release the acquired pool slot. + module AbortableFramer + def abort + stream = @stream + pool = @pool + connection = @connection + @pool = nil + @connection = nil + + begin + io = stream.to_io + io.close unless io.closed? + ensure + pool&.release(connection) + end + end + end + + private_constant :AbortableFramer + + # Configure the native TLS context used by secure WebSockets. + # Peer and hostname verification and HTTP/1.1 ALPN remain SDK-owned. + def initialize( + product_name:, + error_class:, + error_factory:, + dependency_message: nil, + sensitive_query_parameter: nil, + &tls_configurator + ) + @product_name = product_name + @error_class = error_class + @error_factory = error_factory + @dependency_message = dependency_message + @sensitive_query_parameter = sensitive_query_parameter + @tls_configurator = tls_configurator + end + + class Socket + # @api private + def initialize(connection, url:, error_factory:) + @connection = connection + @url = url + @error_factory = error_factory + @aborted = false + end + + def read + @connection.read + rescue StandardError => e + raise @error_factory.call(url: @url, cause: e) + end + + def write(message) + @connection.write(message) + @connection.flush + rescue StandardError => e + raise @error_factory.call(url: @url, cause: e) + end + + def close(code: 1000, reason: "") + @connection.close(code, reason) + rescue StandardError => e + raise @error_factory.call(url: @url, cause: e) + end + + # @api private + def abort + framer = @connection.framer + framer.extend(AbortableFramer) + framer.abort + @aborted = true + rescue StandardError => e + raise @error_factory.call(url: @url, cause: e) + end + + # @api private + def aborted? = @aborted + + def closed? = @aborted || @connection.closed? + end + + def open(url:, headers:, timeout:, **endpoint_options) + load_dependencies(url) + + # Proxy credentials belong only on the CONNECT request assembled from + # proxy configuration. Never forward a caller-supplied value to the + # origin handshake. + headers = headers.reject do |name, _value| + name.to_s.casecmp?("proxy-authorization") + end + + # Classic WebSocket negotiation uses HTTP/1.1. Pinning ALPN also avoids an + # HTTP/2 selection on servers that advertise both protocols. + options = { + alpn_protocols: ::Async::HTTP::Protocol::HTTP11.names, + **endpoint_options + } + if url.scheme == "wss" || @tls_configurator + options[:ssl_context] = build_tls_context(url) + end + + sideband = sensitive_query?(url) + endpoint_url = url.dup + endpoint_url.query = nil if sideband + + endpoint = ::Async::HTTP::Endpoint.parse( + endpoint_url.to_s, + **options + ) + request_target = endpoint.path + request_target = "#{request_target}?#{url.query}" if sideband + proxy_client = nil + if (proxy = proxy_uri(url)) + proxy_endpoint = ::Async::HTTP::Endpoint.parse(proxy_url(proxy).to_s) + proxy_client = ::Async::HTTP::Client.open(proxy_endpoint) + tunnel = ::Async::HTTP::Proxy.new( + proxy_client, + authority(url, include_default_port: true), + trace_safe_headers(proxy_headers(proxy)) + ) + endpoint = tunnel.wrap_endpoint(endpoint) + end + + block_error = nil + # Keep the request timeout scoped to WebSocket negotiation. Endpoint timeouts + # remain installed on the socket and would otherwise terminate healthy idle + # sessions after the ordinary HTTP request timeout. + ::Kernel.Sync() do + client_options = if sideband + { + protocol: TraceSafeProtocol.new( + endpoint.protocol, + sensitive_query_parameter: @sensitive_query_parameter + ) + } + else + {} + end + + client = ::Async::WebSocket::Client.open(endpoint, **client_options) + connection = nil + socket = nil + begin + connection = negotiate( + client, + endpoint, + request_target: request_target, + headers: headers, + timeout: timeout + ) + socket = build_socket(connection, url: url) + begin + yield(socket) + rescue StandardError => e + block_error = e + raise + end + + ensure + close_resources( + connection, + client, + proxy_client, + connection_aborted: socket&.aborted?, + pending_error: $ERROR_INFO + ) + end + end + + rescue StandardError => e + raise if @error_class === e + raise if e.equal?(block_error) + + raise @error_factory.call(url: url, cause: e, http_status: handshake_status(e)) + end + + private def negotiate(client, endpoint, request_target:, headers:, timeout:) + safe_headers = trace_safe_headers(headers) + operation = lambda do + client.connect(authority(endpoint.url), request_target, headers: safe_headers) + end + + return operation.call if timeout.nil? + + ::Async::Task.current.with_timeout(timeout, &operation) + end + + private def build_socket(connection, url:) + Socket.new(connection, url: url, error_factory: @error_factory) + end + + private def sensitive_query?(url) + return false unless @sensitive_query_parameter && url.query + + URI.decode_www_form(url.query).any? { |name, _value| name == @sensitive_query_parameter } + end + + private def build_tls_context(url) + unless url.scheme == "wss" + raise ArgumentError, "TLS configuration requires a wss:// #{@product_name} endpoint" + end + + context = OpenSSL::SSL::SSLContext.new + @tls_configurator&.call(context) + if context.verify_callback + raise ArgumentError, "#{@product_name} WebSocket TLS configuration cannot set verify_callback" + end + + context.set_params(verify_mode: OpenSSL::SSL::VERIFY_PEER) + context.verify_hostname = true + context.alpn_protocols = ::Async::HTTP::Protocol::HTTP11.names + context + end + + private def close_resources(connection, client, proxy_client, connection_aborted:, pending_error:) + cleanup_error = nil + resources = [client, proxy_client] + resources.unshift(connection) unless connection_aborted + resources.compact.each do |resource| + resource.close + rescue StandardError => e + cleanup_error ||= e + end + + raise cleanup_error if pending_error.nil? && cleanup_error + end + + private def load_dependencies(url) + require("async/websocket/client") + require("async/http/endpoint") + require("async/http/proxy") + + rescue LoadError => e + message = @dependency_message || + "#{@product_name} WebSockets require the `async-websocket` gem. " \ + "Add `gem \"async-websocket\"` to your Gemfile." + raise @error_factory.call(url: url, message: message, cause: e) + end + + private def authority(url, include_default_port: false) + host = url.hostname + host = "[#{host}]" if host.include?(":") + default_port = %w[https wss].include?(url.scheme) ? 443 : 80 + return host if !include_default_port && url.port == default_port + + "#{host}:#{url.port}" + end + + private def proxy_uri(url) + policy_url = url.dup + policy_url.scheme = {"ws" => "http", "wss" => "https"}.fetch(url.scheme, url.scheme) + policy_url.find_proxy + end + + private def proxy_url(proxy) + unless %w[http https].include?(proxy.scheme) && proxy.hostname + raise ArgumentError, "#{@product_name} WebSocket proxy must be an absolute HTTP or HTTPS URL" + end + + proxy.dup.tap do |url| + url.user = nil + url.password = nil + end + end + + private def proxy_headers(proxy) + return {} unless proxy.user + + user = URI::RFC2396_PARSER.unescape(proxy.user) + password = URI::RFC2396_PARSER.unescape(proxy.password.to_s) + {"proxy-authorization" => "Basic #{["#{user}:#{password}"].pack("m0")}"} + end + + private def trace_safe_headers(headers) + fields = ::Protocol::HTTP::Headers[headers].to_a + trace_safe_headers_class.new(fields) + end + + private def trace_safe_headers_class + @trace_safe_headers_class ||= Class.new(::Protocol::HTTP::Headers) do + include(TraceSafeHeaders) + end + end + + private def handshake_status(error) + return unless error.is_a?(::Async::WebSocket::ConnectionError) + + error.response.status + end + end + end +end diff --git a/lib/openai/helpers/websocket/client_request.rb b/lib/openai/helpers/websocket/client_request.rb new file mode 100644 index 000000000..e79c167b4 --- /dev/null +++ b/lib/openai/helpers/websocket/client_request.rb @@ -0,0 +1,131 @@ +# frozen_string_literal: true + +module OpenAI + module WebSocket + # Shared handshake request helpers used by product WebSocket facades. + # + # @api private + module ClientRequest + private def with_websocket_connection_retry(error_class:, build:) + request, deadline = build.call(nil) + handshake_completed = false + mark_handshake_completed = -> { handshake_completed = true } + yield(request, mark_handshake_completed) + rescue error_class => e + raise if handshake_completed + raise unless e.http_status == 401 && @workload_identity_auth + + @workload_identity_auth.invalidate_token + refreshed, = build.call(deadline) + yield(refreshed, mark_handshake_completed) + end + + private def shared_websocket_request_options(options, extra_query_message:, max_retries_message:) + opts = options.to_h.dup + OpenAI::RequestOptions.validate!(opts) + extra_query = opts.delete(:extra_query) + unless extra_query.nil? || (extra_query.respond_to?(:empty?) && extra_query.empty?) + raise ArgumentError, extra_query_message + end + + max_retries = opts[:max_retries] + raise ArgumentError, max_retries_message unless max_retries.nil? || max_retries == 0 + + opts + end + + private def parse_shared_websocket_base_url(value, invalid_message:, malformed_message:, preserve_cause:) + return if value.nil? + + uri = URI(value.to_s) + valid_scheme = %w[http https ws wss].include?(uri.scheme) + ambiguous = uri.userinfo || uri.query || uri.fragment + raise ArgumentError, invalid_message unless uri.absolute? && uri.host && valid_scheme && !ambiguous + + uri + rescue URI::Error => e + raise ArgumentError, malformed_message, cause: preserve_cause ? e : nil + end + + private def shared_websocket_url(url) + url = url.dup + url.scheme = {"http" => "ws", "https" => "wss"}.fetch(url.scheme, url.scheme) + url + end + + private def shared_websocket_headers(headers) + headers.except("accept", "content-type").reject do |name, _value| + name.to_s.casecmp?("proxy-authorization") + end + end + + private def build_shared_websocket_connection_request( + path:, + query:, + websocket_base_url:, + options:, + deadline: nil, + validate:, + invalid_base_url_message:, + malformed_base_url_message:, + preserve_base_url_cause:, + extra_query_message:, + max_retries_message:, + timeout_error: + ) + validate.call(websocket_base_url) + websocket_uri = parse_shared_websocket_base_url( + websocket_base_url, + invalid_message: invalid_base_url_message, + malformed_message: malformed_base_url_message, + preserve_cause: preserve_base_url_cause + ) + opts = shared_websocket_request_options( + options, + extra_query_message: extra_query_message, + max_retries_message: max_retries_message + ) + + request = build_request( + {method: :get, path: path, query: query, security: {bearer_auth: true}}, + opts + ) + error_request = websocket_uri ? with_shared_websocket_base_url(request, path:, base_url: websocket_uri) : request + error_url = shared_websocket_url(error_request.fetch(:url)) + + if @workload_identity_auth + deadline ||= request[:timeout]&.then do |timeout| + OpenAI::Internal::Util.monotonic_secs + timeout + end + end + + placeholder = "Bearer #{OpenAI::Client::WORKLOAD_IDENTITY_API_KEY_PLACEHOLDER}" + if @workload_identity_auth && request.fetch(:headers)["authorization"] == placeholder + token = @workload_identity_auth.get_token(deadline: deadline) + request = request.merge(headers: request.fetch(:headers).merge("authorization" => "Bearer #{token}")) + end + + request = prepare_request(request, redirect_count: 0, retry_count: 0) + request = with_shared_websocket_base_url(request, path:, base_url: websocket_uri) if websocket_uri + request = request.merge( + url: shared_websocket_url(request.fetch(:url)), + headers: shared_websocket_headers(request.fetch(:headers)) + ) + request = request_with_remaining_timeout(request, deadline) unless deadline.nil? + [request, deadline] + rescue Timeout::Error => e + error = timeout_error.call(error_url, e) + raise error, cause: error.cause + end + + private def with_shared_websocket_base_url(request, path:, base_url:) + url = OpenAI::Internal::Util.join_parsed_uri( + OpenAI::Internal::Util.parse_uri(base_url.to_s), + {path: OpenAI::Internal::Util.interpolate_path(path)} + ) + url.query = request.fetch(:url).query + request.merge(url: url) + end + end + end +end diff --git a/lib/openai/helpers/websocket/connection.rb b/lib/openai/helpers/websocket/connection.rb new file mode 100644 index 000000000..d1edda261 --- /dev/null +++ b/lib/openai/helpers/websocket/connection.rb @@ -0,0 +1,96 @@ +# frozen_string_literal: true + +module OpenAI + module WebSocket + # Product-neutral socket lifecycle and text I/O. + # + # @api private + class Connection + include Enumerable + + # @return [URI::Generic] + attr_reader :url + + # @api private + def initialize(socket:, url:) + @socket = socket + @url = url + end + + # Yield server events until the remote peer closes the connection. + def each + return enum_for(__method__) unless block_given? + + while (event = receive) + yield(event) + end + + self + end + + # Receive and parse the next server event, or return nil after a clean close. + def receive + data = receive_raw + return nil if data.nil? + + parse_event(data) + end + + # Receive the next raw WebSocket message. + # + # @api private + def receive_raw + message = @socket.read + message&.to_str + end + + # Send an already encoded text message. + # + # @api private + def send_raw(data) + raise connection_error(closed_send_message) if closed? + + text = data.dup + text.force_encoding(Encoding::UTF_8) if text.encoding == Encoding::BINARY + text = text.encode(Encoding::UTF_8) unless text.encoding == Encoding::UTF_8 + raise ArgumentError, invalid_text_message unless text.valid_encoding? + + @socket.write(text) + nil + end + + # Close the connection. + def close(code: 1000, reason: "") + return if closed? + + @socket.close(code: code, reason: reason) + nil + end + + # Abort without waiting for the WebSocket close handshake. + # + # @api private + def abort + return if closed? + + @socket.abort + nil + end + + # @return [Boolean] + def closed? = @socket.closed? + + private def parse_event(_data) + raise NotImplementedError + end + + private def connection_error(_message) + raise NotImplementedError + end + + private def closed_send_message = "Cannot send on a closed WebSocket." + + private def invalid_text_message = "WebSocket text must contain valid UTF-8" + end + end +end diff --git a/lib/openai/helpers/websocket/connection_manager.rb b/lib/openai/helpers/websocket/connection_manager.rb new file mode 100644 index 000000000..762be277b --- /dev/null +++ b/lib/openai/helpers/websocket/connection_manager.rb @@ -0,0 +1,104 @@ +# frozen_string_literal: true + +module OpenAI + module WebSocket + # Shared block-scoped WebSocket opening and cleanup. + # + # @api private + class ConnectionManager + RESERVED_TRANSPORT_OPTIONS = [ + :alpn_protocols, + :headers, + :hostname, + :port, + :protocol, + :scheme, + :ssl_context, + :timeout, + :url + ].freeze + + # @api private + def initialize( + transport:, + transport_options:, + default_transport:, + connection_class:, + request:, + block_error_message:, + abort_after_block:, + transport_error_message: "transport must respond to open", + reserved_options_error_message: nil + ) + @transport = transport + @default_transport = default_transport + @connection_class = connection_class + @request = request + @block_error_message = block_error_message + @abort_after_block = abort_after_block + @transport_error_message = transport_error_message + @reserved_options_error_message = reserved_options_error_message + @transport_options = validated_transport_options(transport_options) + end + + # @api private + def open + raise ArgumentError, @block_error_message unless block_given? + + transport = @transport || @default_transport.call + raise ArgumentError, @transport_error_message unless transport.respond_to?(:open) + + @request.call do |request, mark_handshake_completed| + transport + .open( + url: request.fetch(:url), + headers: request.fetch(:headers), + timeout: request.fetch(:timeout), + **@transport_options + ) do |socket| + mark_handshake_completed.call + connection = @connection_class.new(socket: socket, url: request.fetch(:url)) + begin + yield(connection) + ensure + cleanup(connection) + end + end + end + end + + private def validated_transport_options(transport_options) + options = transport_options.dup.freeze + reserved = options.keys.select do |key| + (key.is_a?(String) || key.is_a?(Symbol)) && RESERVED_TRANSPORT_OPTIONS.include?(key.to_sym) + end + + unless reserved.empty? + message = if @reserved_options_error_message + @reserved_options_error_message.call(reserved) + else + "transport_options cannot include #{reserved.map(&:inspect).join(", ")}" + end + + raise ArgumentError, message + end + + options + end + + private def cleanup(connection) + pending_error = $ERROR_INFO + begin + if @abort_after_block.call(connection, pending_error) + connection.abort unless connection.closed? + else + connection.close unless connection.closed? + end + + rescue StandardError + raise if pending_error.nil? + end + end + end + end +end diff --git a/lib/openai/helpers/websocket/errors.rb b/lib/openai/helpers/websocket/errors.rb new file mode 100644 index 000000000..004aa84cc --- /dev/null +++ b/lib/openai/helpers/websocket/errors.rb @@ -0,0 +1,32 @@ +# frozen_string_literal: true + +module OpenAI + module Errors + # Base class for SDK WebSocket connection failures. + # + # @api private + class WebSocketConnectionError < OpenAI::Errors::Error + attr_reader :url, :http_status + + def cause = @cause.nil? ? super : @cause + + # @api private + def initialize(url:, message: nil, cause: nil, http_status: nil) + @url = sanitized_error_url(url) + @cause = cause + @http_status = http_status + super(message || default_message) + end + + private def default_message = "WebSocket connection error." + + private def sanitized_error_url(url) = url + end + + # Base class for malformed WebSocket messages. + # + # @api private + class WebSocketProtocolError < OpenAI::Errors::Error + end + end +end diff --git a/lib/openai/helpers/websocket/protocol.rb b/lib/openai/helpers/websocket/protocol.rb new file mode 100644 index 000000000..7c56285c1 --- /dev/null +++ b/lib/openai/helpers/websocket/protocol.rb @@ -0,0 +1,26 @@ +# frozen_string_literal: true + +module OpenAI + module WebSocket + # Shared helpers for discriminated JSON WebSocket events. + # + # @api private + module Protocol + private def discriminator_values(union) + union.variants.to_h do |variant| + value = variant.fields.fetch(:type).fetch(:const) + [value.to_s, true] + end + end + + private def event_type(event, message:) + raise ArgumentError, message unless event.is_a?(Hash) + + type = event[:type] + return type if type.is_a?(String) || type.is_a?(Symbol) + + raise ArgumentError, message + end + end + end +end diff --git a/rbi/openai/helpers/realtime/extensions.rbi b/rbi/openai/helpers/realtime/extensions.rbi index 409944a3a..82ad5fe57 100644 --- a/rbi/openai/helpers/realtime/extensions.rbi +++ b/rbi/openai/helpers/realtime/extensions.rbi @@ -48,7 +48,7 @@ module OpenAI end module Errors - class RealtimeConnectionError < OpenAI::Errors::Error + class RealtimeConnectionError < OpenAI::Errors::WebSocketConnectionError sig { returns(URI::Generic) } attr_reader :url @@ -74,7 +74,7 @@ module OpenAI end end - class RealtimeProtocolError < OpenAI::Errors::Error + class RealtimeProtocolError < OpenAI::Errors::WebSocketProtocolError sig { returns(String) } attr_reader :data diff --git a/rbi/openai/helpers/responses_websocket/connection.rbi b/rbi/openai/helpers/responses_websocket/connection.rbi new file mode 100644 index 000000000..c9908b959 --- /dev/null +++ b/rbi/openai/helpers/responses_websocket/connection.rbi @@ -0,0 +1,82 @@ +# typed: strong + +module OpenAI + module Models + module Responses + class Connection + include Enumerable + + ServerEvent = T.type_alias do + T.any( + OpenAI::Responses::ResponsesServerEvent::Variants, + OpenAI::Responses::UnknownServerEvent + ) + end + + ClientEvent = T.type_alias do + T.any( + OpenAI::Responses::ResponsesClientEvent::Variants, + OpenAI::Internal::AnyHash + ) + end + + Elem = type_member { {fixed: ServerEvent} } + + sig { returns(URI::Generic) } + attr_reader :url + + sig { returns(OpenAI::Responses::ConnectionResources::Response) } + attr_reader :response + + # @api private + sig { params(socket: T.anything, url: URI::Generic).returns(T.attached_class) } + def self.new(socket:, url:) + end + + sig do + params(block: T.nilable(T.proc.params(event: ServerEvent).void)).returns( + T.any(OpenAI::Responses::Connection, T::Enumerator[ServerEvent]) + ) + end + def each(&block) + end + + sig { returns(T.nilable(ServerEvent)) } + def receive + end + + # @api private + sig { returns(T.nilable(String)) } + def receive_raw + end + + sig { params(event: ClientEvent).void } + def send_event(event) + end + + # @api private + sig { params(data: String).void } + def send_raw(data) + end + + sig { params(code: Integer, reason: String).void } + def close(code: 1000, reason: "") + end + + # @api private + sig { void } + def abort + end + + sig { returns(T::Boolean) } + def closed? + end + + # @api private + sig { returns(T::Boolean) } + def poisoned? + end + end + end + end +end diff --git a/rbi/openai/helpers/responses_websocket/connection_manager.rbi b/rbi/openai/helpers/responses_websocket/connection_manager.rbi new file mode 100644 index 000000000..38c6faae4 --- /dev/null +++ b/rbi/openai/helpers/responses_websocket/connection_manager.rbi @@ -0,0 +1,39 @@ +# typed: strong + +module OpenAI + module Models + module Responses + class ConnectionManager + # @api private + sig do + params( + client: OpenAI::Client, + websocket_base_url: T.nilable(String), + transport: T.anything, + request_options: T.nilable(OpenAI::RequestOptions::OrHash), + transport_options: T::Hash[Symbol, T.anything] + ) + .returns(T.attached_class) + end + def self.new( + client:, + websocket_base_url:, + transport:, + request_options:, + transport_options: + ) + end + + # @api private + sig do + params( + block: T.proc.params(connection: OpenAI::Responses::Connection).returns(T.anything) + ) + .returns(T.anything) + end + def open(&block) + end + end + end + end +end diff --git a/rbi/openai/helpers/responses_websocket/connection_resources.rbi b/rbi/openai/helpers/responses_websocket/connection_resources.rbi new file mode 100644 index 000000000..59464627b --- /dev/null +++ b/rbi/openai/helpers/responses_websocket/connection_resources.rbi @@ -0,0 +1,124 @@ +# typed: strong + +module OpenAI + module Models + module Responses + module ConnectionResources + class Response + # @api private + sig { params(connection: OpenAI::Responses::Connection).returns(T.attached_class) } + def self.new(connection) + end + + sig do + params( + background: T.nilable(T::Boolean), + context_management: T.nilable( + T::Array[OpenAI::Responses::ResponsesClientEvent::ResponseCreate::ContextManagement::OrHash] + ), + conversation: T.nilable(T.any(String, OpenAI::Responses::ResponseConversationParam::OrHash)), + include: T.nilable(T::Array[OpenAI::Responses::ResponseIncludable::OrSymbol]), + input: OpenAI::Responses::ResponsesClientEvent::ResponseCreate::Input::Variants, + instructions: T.nilable(String), + max_output_tokens: T.nilable(Integer), + max_tool_calls: T.nilable(Integer), + metadata: T.nilable(T::Hash[Symbol, String]), + model: T.any(String, OpenAI::ChatModel::OrSymbol, OpenAI::ResponsesModel::ResponsesOnlyModel::OrSymbol), + moderation: T.nilable(OpenAI::Responses::ResponsesClientEvent::ResponseCreate::Moderation::OrHash), + parallel_tool_calls: T.nilable(T::Boolean), + previous_response_id: T.nilable(String), + prompt: T.nilable(OpenAI::Responses::ResponsePrompt::OrHash), + prompt_cache_key: T.nilable(String), + prompt_cache_options: OpenAI::Responses::ResponsesClientEvent::ResponseCreate::PromptCacheOptions::OrHash, + prompt_cache_retention: T.nilable( + OpenAI::Responses::ResponsesClientEvent::ResponseCreate::PromptCacheRetention::OrSymbol + ), + reasoning: T.nilable(OpenAI::Reasoning::OrHash), + safety_identifier: T.nilable(String), + service_tier: T.nilable(OpenAI::Responses::ResponsesClientEvent::ResponseCreate::ServiceTier::OrSymbol), + store: T.nilable(T::Boolean), + stream: T.nilable(T::Boolean), + stream_id: String, + stream_options: T.nilable(OpenAI::Responses::ResponsesClientEvent::ResponseCreate::StreamOptions::OrHash), + temperature: T.nilable(Float), + text: OpenAI::Responses::ResponseTextConfig::OrHash, + tool_choice: T.any( + OpenAI::Responses::ToolChoiceOptions::OrSymbol, + OpenAI::Responses::ToolChoiceAllowed::OrHash, + OpenAI::Responses::ToolChoiceTypes::OrHash, + OpenAI::Responses::ToolChoiceFunction::OrHash, + OpenAI::Responses::ToolChoiceMcp::OrHash, + OpenAI::Responses::ToolChoiceCustom::OrHash, + OpenAI::Responses::ResponsesClientEvent::ResponseCreate::ToolChoice::SpecificProgrammaticToolCallingParam::OrHash, + OpenAI::Responses::ToolChoiceApplyPatch::OrHash, + OpenAI::Responses::ToolChoiceShell::OrHash + ), + tools: T::Array[ + T.any( + OpenAI::Responses::FunctionTool::OrHash, + OpenAI::Responses::FileSearchTool::OrHash, + OpenAI::Responses::ComputerTool::OrHash, + OpenAI::Responses::ComputerUsePreviewTool::OrHash, + OpenAI::Responses::Tool::Mcp::OrHash, + OpenAI::Responses::Tool::CodeInterpreter::OrHash, + OpenAI::Responses::Tool::ProgrammaticToolCalling::OrHash, + OpenAI::Responses::Tool::ImageGeneration::OrHash, + OpenAI::Responses::Tool::LocalShell::OrHash, + OpenAI::Responses::FunctionShellTool::OrHash, + OpenAI::Responses::CustomTool::OrHash, + OpenAI::Responses::NamespaceTool::OrHash, + OpenAI::Responses::ToolSearchTool::OrHash, + OpenAI::Responses::ApplyPatchTool::OrHash, + OpenAI::Responses::WebSearchTool::OrHash, + OpenAI::Responses::WebSearchPreviewTool::OrHash + ) + ], + top_logprobs: T.nilable(Integer), + top_p: T.nilable(Float), + truncation: T.nilable(OpenAI::Responses::ResponsesClientEvent::ResponseCreate::Truncation::OrSymbol), + user: String, + params: T::Hash[T.any(String, Symbol), T.anything] + ) + .void + end + def create( + background: nil, + context_management: nil, + conversation: nil, + include: nil, + input: nil, + instructions: nil, + max_output_tokens: nil, + max_tool_calls: nil, + metadata: nil, + model: nil, + moderation: nil, + parallel_tool_calls: nil, + previous_response_id: nil, + prompt: nil, + prompt_cache_key: nil, + prompt_cache_options: nil, + prompt_cache_retention: nil, + reasoning: nil, + safety_identifier: nil, + service_tier: nil, + store: nil, + stream: nil, + stream_id: nil, + stream_options: nil, + temperature: nil, + text: nil, + tool_choice: nil, + tools: nil, + top_logprobs: nil, + top_p: nil, + truncation: nil, + user: nil, + **params + ) + end + end + end + end + end +end diff --git a/rbi/openai/helpers/responses_websocket/extensions.rbi b/rbi/openai/helpers/responses_websocket/extensions.rbi new file mode 100644 index 000000000..0af3526a5 --- /dev/null +++ b/rbi/openai/helpers/responses_websocket/extensions.rbi @@ -0,0 +1,108 @@ +# typed: strong + +module OpenAI + class Client + # @api private + sig do + params( + websocket_base_url: T.nilable(String), + options: T.nilable(OpenAI::RequestOptions::OrHash), + block: T + .proc + .params( + request: OpenAI::Internal::Transport::BaseClient::RequestInput, + mark_handshake_completed: T.proc.void + ) + .returns(T.anything) + ) + .returns(T.anything) + end + def with_responses_websocket_connection_request( + websocket_base_url: nil, + options: nil, + &block + ) + end + end + + module Errors + class ResponsesConnectionError < OpenAI::Errors::WebSocketConnectionError + sig { returns(URI::Generic) } + attr_reader :url + + sig { returns(T.nilable(Integer)) } + attr_reader :http_status + + sig { returns(T.nilable(Exception)) } + def cause + end + + # @api private + sig do + params( + url: URI::Generic, + message: T.nilable(String), + http_status: T.nilable(Integer) + ) + .returns(T.attached_class) + end + def self.new(url:, message: nil, http_status: nil) + end + end + + class ResponsesProtocolError < OpenAI::Errors::WebSocketProtocolError + sig { returns(T.nilable(Exception)) } + def cause + end + + # @api private + sig { returns(T.attached_class) } + def self.new + end + end + + class ResponsesClientEventError < OpenAI::Errors::Error + sig { returns(T.nilable(Exception)) } + def cause + end + + # @api private + sig { returns(T.attached_class) } + def self.new + end + end + + class ResponsesSendError < ResponsesConnectionError + sig { returns(Symbol) } + attr_reader :outcome + + # @api private + sig { params(url: URI::Generic).returns(T.attached_class) } + def self.new(url:) + end + end + end + + module Resources + class Responses + sig do + params( + websocket_base_url: T.nilable(String), + request_options: T.nilable(OpenAI::RequestOptions::OrHash), + transport: T.anything, + transport_options: T::Hash[Symbol, T.anything], + block: T.proc.params(connection: OpenAI::Responses::Connection).returns(T.anything) + ) + .returns(T.anything) + end + def connect( + websocket_base_url: nil, + request_options: nil, + transport: nil, + transport_options: {}, + &block + ) + end + end + end +end diff --git a/rbi/openai/helpers/responses_websocket/transports/async_websocket.rbi b/rbi/openai/helpers/responses_websocket/transports/async_websocket.rbi new file mode 100644 index 000000000..2900e5de3 --- /dev/null +++ b/rbi/openai/helpers/responses_websocket/transports/async_websocket.rbi @@ -0,0 +1,28 @@ +# typed: strong + +module OpenAI + module Models + module Responses + module Transports + class AsyncWebSocket + sig { void } + def initialize + end + + sig do + params( + url: URI::Generic, + headers: T::Hash[String, String], + timeout: T.nilable(Float), + options: T.anything, + block: T.proc.params(socket: T.anything).returns(T.anything) + ) + .returns(T.anything) + end + def open(url:, headers:, timeout:, **options, &block) + end + end + end + end + end +end diff --git a/rbi/openai/helpers/responses_websocket/unknown_server_event.rbi b/rbi/openai/helpers/responses_websocket/unknown_server_event.rbi new file mode 100644 index 000000000..ee0ad4814 --- /dev/null +++ b/rbi/openai/helpers/responses_websocket/unknown_server_event.rbi @@ -0,0 +1,29 @@ +# typed: strong + +module OpenAI + module Models + module Responses + class UnknownServerEvent + # @api private + sig do + params(data: T::Hash[Symbol, T.anything]).returns(T.attached_class) + end + def self.new(data:) + end + + sig { returns(Symbol) } + attr_reader :type + + sig { returns(T::Hash[Symbol, T.anything]) } + attr_reader :data + + sig { returns(T.anything) } + attr_reader :stream_id + + sig { returns(T::Hash[Symbol, T.anything]) } + def to_h + end + end + end + end +end diff --git a/rbi/openai/helpers/websocket.rbi b/rbi/openai/helpers/websocket.rbi new file mode 100644 index 000000000..6baddffc1 --- /dev/null +++ b/rbi/openai/helpers/websocket.rbi @@ -0,0 +1,33 @@ +# typed: strong + +module OpenAI + module Errors + class WebSocketConnectionError < OpenAI::Errors::Error + sig { returns(URI::Generic) } + attr_reader :url + + sig { returns(T.nilable(Integer)) } + attr_reader :http_status + + sig { returns(T.nilable(Exception)) } + def cause + end + + # @api private + sig do + params( + url: URI::Generic, + message: T.nilable(String), + cause: T.nilable(Exception), + http_status: T.nilable(Integer) + ) + .returns(T.attached_class) + end + def self.new(url:, message: nil, cause: nil, http_status: nil) + end + end + + class WebSocketProtocolError < OpenAI::Errors::Error + end + end +end diff --git a/sig/openai/helpers/realtime/extensions.rbs b/sig/openai/helpers/realtime/extensions.rbs index ea705230b..a55f6e0a3 100644 --- a/sig/openai/helpers/realtime/extensions.rbs +++ b/sig/openai/helpers/realtime/extensions.rbs @@ -23,7 +23,7 @@ module OpenAI end module Errors - class RealtimeConnectionError < OpenAI::Errors::Error + class RealtimeConnectionError < OpenAI::Errors::WebSocketConnectionError attr_reader url: URI::Generic # @api private @@ -39,7 +39,7 @@ module OpenAI ) -> void end - class RealtimeProtocolError < OpenAI::Errors::Error + class RealtimeProtocolError < OpenAI::Errors::WebSocketProtocolError attr_reader data: String def cause: -> StandardError? diff --git a/sig/openai/helpers/responses_websocket/connection.rbs b/sig/openai/helpers/responses_websocket/connection.rbs new file mode 100644 index 000000000..3df640d39 --- /dev/null +++ b/sig/openai/helpers/responses_websocket/connection.rbs @@ -0,0 +1,48 @@ +module OpenAI + module Models + module Responses + type connection_server_event = + OpenAI::Models::Responses::responses_server_event + | OpenAI::Responses::UnknownServerEvent + + type connection_client_event = + OpenAI::Models::Responses::responses_client_event + | ::Hash[Symbol | String, untyped] + + class Connection + include Enumerable[OpenAI::Models::Responses::connection_server_event] + + attr_reader url: URI::Generic + attr_reader response: OpenAI::Responses::ConnectionResources::Response + + # @api private + def initialize: (socket: top, url: URI::Generic) -> void + def each: + -> Enumerator[OpenAI::Models::Responses::connection_server_event, self] + | { + (OpenAI::Models::Responses::connection_server_event event) -> void + } -> self + def receive: -> OpenAI::Models::Responses::connection_server_event? + + # @api private + def receive_raw: -> String? + + def send_event: ( + OpenAI::Models::Responses::connection_client_event event + ) -> nil + + # @api private + def send_raw: (String data) -> nil + + def close: (?code: Integer, ?reason: String) -> nil + + # @api private + def abort: -> nil + def closed?: -> bool + + # @api private + def poisoned?: -> bool + end + end + end +end diff --git a/sig/openai/helpers/responses_websocket/connection_manager.rbs b/sig/openai/helpers/responses_websocket/connection_manager.rbs new file mode 100644 index 000000000..c79dfc922 --- /dev/null +++ b/sig/openai/helpers/responses_websocket/connection_manager.rbs @@ -0,0 +1,19 @@ +module OpenAI + module Models + module Responses + class ConnectionManager + # @api private + def initialize: ( + client: OpenAI::Client, + websocket_base_url: String?, + transport: top, + request_options: OpenAI::request_opts?, + transport_options: ::Hash[Symbol, top] + ) -> void + + # @api private + def open: { (OpenAI::Responses::Connection connection) -> top } -> top + end + end + end +end diff --git a/sig/openai/helpers/responses_websocket/connection_resources.rbs b/sig/openai/helpers/responses_websocket/connection_resources.rbs new file mode 100644 index 000000000..62a029451 --- /dev/null +++ b/sig/openai/helpers/responses_websocket/connection_resources.rbs @@ -0,0 +1,47 @@ +module OpenAI + module Models + module Responses + module ConnectionResources + class Response + # @api private + def initialize: (OpenAI::Responses::Connection connection) -> void + def create: ( + ?background: bool?, + ?context_management: ::Array[OpenAI::Responses::ResponsesClientEvent::ResponseCreate::ContextManagement]?, + ?conversation: OpenAI::Models::Responses::ResponsesClientEvent::ResponseCreate::conversation?, + ?include: ::Array[OpenAI::Models::Responses::response_includable]?, + ?input: OpenAI::Models::Responses::ResponsesClientEvent::ResponseCreate::input, + ?instructions: String?, + ?max_output_tokens: Integer?, + ?max_tool_calls: Integer?, + ?metadata: OpenAI::Models::metadata?, + ?model: OpenAI::Models::responses_model, + ?moderation: OpenAI::Responses::ResponsesClientEvent::ResponseCreate::Moderation?, + ?parallel_tool_calls: bool?, + ?previous_response_id: String?, + ?prompt: OpenAI::Responses::ResponsePrompt?, + ?prompt_cache_key: String?, + ?prompt_cache_options: OpenAI::Responses::ResponsesClientEvent::ResponseCreate::PromptCacheOptions, + ?prompt_cache_retention: OpenAI::Models::Responses::ResponsesClientEvent::ResponseCreate::prompt_cache_retention?, + ?reasoning: OpenAI::Reasoning?, + ?safety_identifier: String?, + ?service_tier: OpenAI::Models::Responses::ResponsesClientEvent::ResponseCreate::service_tier?, + ?store: bool?, + ?stream: bool?, + ?stream_id: String, + ?stream_options: OpenAI::Responses::ResponsesClientEvent::ResponseCreate::StreamOptions?, + ?temperature: Float?, + ?text: OpenAI::Responses::ResponseTextConfig, + ?tool_choice: OpenAI::Models::Responses::ResponsesClientEvent::ResponseCreate::tool_choice, + ?tools: ::Array[OpenAI::Models::Responses::tool], + ?top_logprobs: Integer?, + ?top_p: Float?, + ?truncation: OpenAI::Models::Responses::ResponsesClientEvent::ResponseCreate::truncation?, + ?user: String, + **top params + ) -> nil + end + end + end + end +end diff --git a/sig/openai/helpers/responses_websocket/extensions.rbs b/sig/openai/helpers/responses_websocket/extensions.rbs new file mode 100644 index 000000000..14a67c914 --- /dev/null +++ b/sig/openai/helpers/responses_websocket/extensions.rbs @@ -0,0 +1,55 @@ +module OpenAI + class Client + # @api private + def with_responses_websocket_connection_request: ( + ?websocket_base_url: String?, + ?options: OpenAI::request_opts? + ) { + ( + OpenAI::Internal::Transport::BaseClient::request_input request, + ^-> void mark_handshake_completed + ) -> top + } -> top + end + + module Errors + class ResponsesConnectionError < OpenAI::Errors::WebSocketConnectionError + attr_reader url: URI::Generic + attr_reader http_status: Integer? + def cause: -> Exception? + def initialize: ( + url: URI::Generic, + ?message: String?, + ?http_status: Integer? + ) -> void + end + + class ResponsesProtocolError < OpenAI::Errors::WebSocketProtocolError + def cause: -> Exception? + def initialize: -> void + end + + class ResponsesClientEventError < OpenAI::Errors::Error + def cause: -> Exception? + def initialize: -> void + end + + class ResponsesSendError < ResponsesConnectionError + attr_reader outcome: Symbol + def initialize: (url: URI::Generic) -> void + end + end + + module Resources + class Responses + def connect: ( + ?websocket_base_url: String?, + ?request_options: OpenAI::request_opts?, + ?transport: top?, + ?transport_options: ::Hash[Symbol, top] + ) { + (OpenAI::Responses::Connection connection) -> top + } -> top + end + end +end diff --git a/sig/openai/helpers/responses_websocket/transports/async_websocket.rbs b/sig/openai/helpers/responses_websocket/transports/async_websocket.rbs new file mode 100644 index 000000000..70bd21ae2 --- /dev/null +++ b/sig/openai/helpers/responses_websocket/transports/async_websocket.rbs @@ -0,0 +1,19 @@ +module OpenAI + module Models + module Responses + module Transports + class AsyncWebSocket + def initialize: -> void + def open: ( + url: URI::Generic, + headers: ::Hash[String, String], + timeout: Float?, + **top options + ) { + (top socket) -> top + } -> top + end + end + end + end +end diff --git a/sig/openai/helpers/responses_websocket/unknown_server_event.rbs b/sig/openai/helpers/responses_websocket/unknown_server_event.rbs new file mode 100644 index 000000000..04d5ea57f --- /dev/null +++ b/sig/openai/helpers/responses_websocket/unknown_server_event.rbs @@ -0,0 +1,15 @@ +module OpenAI + module Models + module Responses + class UnknownServerEvent + # @api private + def initialize: (data: ::Hash[Symbol, top]) -> void + + attr_reader type: Symbol + attr_reader data: ::Hash[Symbol, top] + attr_reader stream_id: top + def to_h: -> ::Hash[Symbol, top] + end + end + end +end diff --git a/sig/openai/helpers/websocket.rbs b/sig/openai/helpers/websocket.rbs new file mode 100644 index 000000000..2b7b8874c --- /dev/null +++ b/sig/openai/helpers/websocket.rbs @@ -0,0 +1,18 @@ +module OpenAI + module Errors + class WebSocketConnectionError < OpenAI::Errors::Error + attr_reader url: URI::Generic + attr_reader http_status: Integer? + def cause: -> Exception? + def initialize: ( + url: URI::Generic, + ?message: String?, + ?cause: Exception?, + ?http_status: Integer? + ) -> void + end + + class WebSocketProtocolError < OpenAI::Errors::Error + end + end +end diff --git a/test/openai/realtime/async_websocket_transport_test.rb b/test/openai/realtime/async_websocket_transport_test.rb index b40b2694a..af4d0fcfb 100644 --- a/test/openai/realtime/async_websocket_transport_test.rb +++ b/test/openai/realtime/async_websocket_transport_test.rb @@ -158,12 +158,18 @@ def close = @closed = true url = URI("wss://example.com/v1/realtime?model=gpt-realtime-2.1") + socket = nil Async::HTTP::Endpoint.stub(:parse, parser) do Async::WebSocket::Client.stub(:open, client) do - transport.open(url: url, headers: {}, timeout: nil) { |_socket| nil } + transport.open(url: url, headers: {}, timeout: nil) { |value| socket = value } end end + assert_instance_of(OpenAI::Realtime::Transports::AsyncWebSocket::Socket, socket) + assert_instance_of( + OpenAI::Realtime::Transports::AsyncWebSocket::Socket, + OpenAI::Realtime::Transports::AsyncWebSocket::Socket.new(connection, url: url) + ) assert_equal(OpenSSL::SSL::VERIFY_PEER, parsed_endpoint.ssl_context.verify_mode) assert_predicate(parsed_endpoint.ssl_context, :verify_hostname) assert_predicate(client, :closed) diff --git a/test/openai/responses_websocket/connection_test.rb b/test/openai/responses_websocket/connection_test.rb new file mode 100644 index 000000000..ec14310b1 --- /dev/null +++ b/test/openai/responses_websocket/connection_test.rb @@ -0,0 +1,648 @@ +# frozen_string_literal: true + +require_relative "connection_test_support" + +class OpenAI::Test::ResponsesWebSocketConnectionTest < Minitest::Test + include OpenAI::Test::ResponsesWebSocketConnectionTestSupport + + class SerializerMetadataProbe < OpenAI::Internal::Type::BaseModel + required :type, const: :future + optional :ruby_name, String, api_name: :apiName + end + + def test_connect_opens_a_block_scoped_connection_and_sends_response_create + socket = FakeSocket.new(text_delta("hello", stream_id: "turn_1")) + transport = FakeTransport.new(socket) + event = nil + + result = client + .responses + .connect( + request_options: {extra_headers: {"X-Trace-ID" => "trace_1"}, timeout: 12}, + transport: transport, + transport_options: {max_frame_size: 1_024} + ) do |connection| + assert_instance_of(OpenAI::Responses::Connection, connection) + assert_nil(connection.response.create(model: "gpt-5.2", input: "hi", stream_id: "turn_1")) + event = connection.receive + :block_result + end + + assert_equal(:block_result, result) + assert_instance_of(OpenAI::Responses::ResponsesServerEvent::ResponseTextWsDelta, event) + assert_equal("hello", event.delta) + assert_equal("turn_1", event.stream_id) + assert_equal("wss://example.com/v1/responses", transport.open_args.fetch(:url).to_s) + assert_equal("Bearer test-key", transport.open_args.fetch(:headers).fetch("authorization")) + assert_equal("trace_1", transport.open_args.fetch(:headers).fetch("x-trace-id")) + assert_equal(12.0, transport.open_args.fetch(:timeout)) + assert_equal({max_frame_size: 1_024}, transport.open_args.fetch(:options)) + assert_equal( + {"type" => "response.create", "model" => "gpt-5.2", "input" => "hi", "stream_id" => "turn_1"}, + JSON.parse(socket.writes.fetch(0)) + ) + assert_predicate(socket, :closed?) + end + + def test_connect_requires_a_block + error = assert_raises(ArgumentError) { client.responses.connect } + + assert_equal("A block is required to open a Responses WebSocket.", error.message) + end + + def test_raw_io_helpers_remain_public_private_sdk_plumbing + socket = FakeSocket.new("raw server message") + + client.responses.connect(transport: FakeTransport.new(socket)) do |connection| + assert_equal("raw server message", connection.receive_raw) + assert_nil(connection.send_raw(JSON.generate(type: "response.create"))) + end + + assert_equal({"type" => "response.create"}, JSON.parse(socket.writes.fetch(0))) + end + + def test_response_create_keeps_its_fixed_discriminator + socket = FakeSocket.new + + client.responses.connect(transport: FakeTransport.new(socket)) do |connection| + connection.response.create(type: :"response.cancel", model: "gpt-5.2") + connection.response.create(**{"type" => :"response.cancel", :model => "gpt-5.2"}) + end + + assert_equal( + {"type" => "response.create", "model" => "gpt-5.2"}, + JSON.parse(socket.writes.fetch(0)) + ) + assert_equal( + {"type" => "response.create", "model" => "gpt-5.2"}, + JSON.parse(socket.writes.fetch(1)) + ) + end + + def test_unknown_event_is_observable_without_exposing_payload_in_inspect + secret = "secret-event-payload" + socket = FakeSocket.new(JSON.generate(type: "response.future", stream_id: "turn_1", secret: secret)) + event = nil + + client.responses.connect(transport: FakeTransport.new(socket)) { |connection| event = connection.receive } + + assert_instance_of(OpenAI::Responses::UnknownServerEvent, event) + assert_equal(:"response.future", event.type) + assert_equal(secret, event.data.fetch(:secret)) + refute_includes(event.inspect, secret) + end + + def test_function_call_output_item_without_computed_parsed_is_received + socket = FakeSocket.new( + JSON.generate( + type: "response.output_item.added", + sequence_number: 1, + output_index: 0, + item: { + type: "function_call", + id: "fc_1", + call_id: "call_1", + name: "lookup", + arguments: "{}" + } + ) + ) + event = nil + + client.responses.connect(transport: FakeTransport.new(socket)) { |connection| event = connection.receive } + + assert_instance_of(OpenAI::Responses::ResponsesServerEvent::ResponseOutputItemWsAdded, event) + assert_instance_of(OpenAI::Responses::ResponseFunctionToolCall, event.item) + assert_equal("call_1", event.item.call_id) + end + + def test_replayed_function_call_without_computed_parsed_is_sent + socket = FakeSocket.new + + client.responses.connect(transport: FakeTransport.new(socket)) do |connection| + connection.response.create( + model: "gpt-5.2", + input: [ + { + type: "function_call", + id: "fc_1", + call_id: "call_1", + name: "lookup", + arguments: "{}" + } + ] + ) + end + + item = JSON.parse(socket.writes.fetch(0)).fetch("input").fetch(0) + assert_equal("function_call", item.fetch("type")) + assert_equal("call_1", item.fetch("call_id")) + refute(item.key?("parsed")) + end + + def test_partial_known_event_is_best_effort_typed_without_poisoning + socket = FakeSocket.new( + JSON.generate( + type: "response.content_part.added", + sequence_number: 1, + item_id: "item_1", + output_index: 0, + content_index: 0, + part: {type: "output_text", text: ""} + ), + text_delta("still-readable") + ) + + client.responses.connect(transport: FakeTransport.new(socket)) do |connection| + event = connection.receive + assert_instance_of(OpenAI::Responses::ResponsesServerEvent::ResponseContentPartWsAdded, event) + assert_equal("still-readable", connection.receive.delta) + end + end + + def test_parse_failures_are_payload_free + secret = "secret-invalid-json" + socket = FakeSocket.new("{#{secret}") + + error = assert_raises(OpenAI::Errors::ResponsesProtocolError) do + client.responses.connect(transport: FakeTransport.new(socket)) { |connection| connection.receive } + end + + assert_equal("Invalid Responses WebSocket event.", error.message) + refute_includes(error.full_message, secret) + assert_nil(error.cause) + end + + def test_failed_write_poisons_connection_and_reports_unknown_outcome + socket = FailingWriteSocket.new + connection = nil + + error = assert_raises(OpenAI::Errors::ResponsesSendError) do + client.responses.connect(transport: FakeTransport.new(socket)) do |opened| + connection = opened + opened.response.create(model: "gpt-5.2", input: "sensitive-body") + end + end + + assert_equal(:unknown, error.outcome) + assert_nil(error.cause) + assert_predicate(socket, :aborted?) + assert_raises(OpenAI::Errors::ResponsesConnectionError) do + connection.response.create(model: "gpt-5.2", input: "again") + end + end + + def test_close_aborts_after_ambiguous_write + socket = FailingWriteSocket.new + + client.responses.connect(transport: FakeTransport.new(socket)) do |connection| + assert_raises(OpenAI::Errors::ResponsesSendError) do + connection.response.create(model: "gpt-5.2") + end + + assert_nil(connection.close) + end + + assert_predicate(socket, :aborted?) + assert_nil(socket.close_args) + end + + def test_nested_generated_models_keep_serializer_metadata + socket = FakeSocket.new + nested = SerializerMetadataProbe.new(ruby_name: "kept") + + client.responses.connect(transport: FakeTransport.new(socket)) do |connection| + connection.send_event(type: "response.create", future: nested) + end + + assert_equal( + {"type" => "future", "apiName" => "kept"}, + JSON.parse(socket.writes.fetch(0)).fetch("future") + ) + end + + def test_nested_generated_model_rejects_api_name_collision_before_write + socket = FakeSocket.new + nested = SerializerMetadataProbe.new(ruby_name: "good", apiName: "bad") + + assert_raises(OpenAI::Errors::ResponsesClientEventError) do + client.responses.connect(transport: FakeTransport.new(socket)) do |connection| + connection.send_event(type: "response.create", future: nested) + end + end + + assert_empty(socket.writes) + end + + def test_raw_hash_without_type_is_forwarded + socket = FakeSocket.new + + client.responses.connect(transport: FakeTransport.new(socket)) do |connection| + connection.send_event(model: "gpt-5.2") + end + + assert_equal({"model" => "gpt-5.2"}, JSON.parse(socket.writes.fetch(0))) + end + + def test_nested_mixed_keys_are_rejected_before_write + socket = FakeSocket.new + + assert_raises(OpenAI::Errors::ResponsesClientEventError) do + client.responses.connect(transport: FakeTransport.new(socket)) do |connection| + connection.send_event( + type: "response.create", + model: "gpt-5.2", + prompt_cache_options: {:mode => "in_memory", "mode" => "24h"} + ) + end + end + + assert_empty(socket.writes) + end + + def test_non_string_json_object_keys_are_rejected_before_write + socket = FakeSocket.new + + assert_raises(OpenAI::Errors::ResponsesClientEventError) do + client.responses.connect(transport: FakeTransport.new(socket)) do |connection| + connection.response.create(model: "gpt-5.2", future: {1 => "first", "1" => "second"}) + end + end + + assert_empty(socket.writes) + end + + def test_newer_fields_are_forwarded + socket = FakeSocket.new + + client.responses.connect(transport: FakeTransport.new(socket)) do |connection| + connection.response.create(model: "gpt-5.2", multi_agent: true, stream_id: 123) + end + + payload = JSON.parse(socket.writes.fetch(0)) + assert_equal(true, payload.fetch("multi_agent")) + assert_equal(123, payload.fetch("stream_id")) + end + + def test_opaque_input_maps_can_use_beta_looking_names + socket = FakeSocket.new + + client.responses.connect(transport: FakeTransport.new(socket)) do |connection| + assert_nil( + connection.response.create( + model: "gpt-5.2", + input: [ + { + type: "tool_search_call", + arguments: {type: "agent_message", agent: "opaque-value"} + } + ] + ) + ) + end + + assert_equal( + "opaque-value", + JSON.parse(socket.writes.fetch(0)).dig("input", 0, "arguments", "agent") + ) + end + + def test_opaque_maps_can_use_agent_keys_without_becoming_beta_events + socket = FakeSocket.new + + client.responses.connect(transport: FakeTransport.new(socket)) do |connection| + assert_nil(connection.response.create(model: "gpt-5.2", metadata: {agent: "ordinary-tag"})) + end + + assert_equal("ordinary-tag", JSON.parse(socket.writes.fetch(0)).dig("metadata", "agent")) + end + + def test_cyclic_client_event_data_is_rejected_without_writing + socket = FakeSocket.new + cycle = {} + cycle[:self] = cycle + + error = assert_raises(OpenAI::Errors::ResponsesClientEventError) do + client.responses.connect(transport: FakeTransport.new(socket)) do |connection| + connection.response.create(model: "gpt-5.2", metadata: cycle) + end + end + + assert_equal("Invalid Responses WebSocket client event.", error.message) + assert_empty(socket.writes) + end + + def test_cyclic_typed_client_event_data_is_rejected_without_writing + socket = FakeSocket.new + cycle = {} + cycle[:self] = cycle + event = OpenAI::Responses::ResponsesClientEvent::ResponseCreate.new( + type: :"response.create", + model: "gpt-5.2", + future: cycle + ) + + assert_raises(OpenAI::Errors::ResponsesClientEventError) do + client.responses.connect(transport: FakeTransport.new(socket)) do |connection| + connection.send_event(event) + end + end + + assert_empty(socket.writes) + end + + def test_each_owns_the_read_lease_and_rejects_nested_receive + socket = FakeSocket.new(text_delta("hello")) + + error = assert_raises(OpenAI::Errors::ResponsesConnectionError) do + client.responses.connect(transport: FakeTransport.new(socket)) do |connection| + connection.each { |_event| connection.receive } + end + end + + assert_equal("Responses WebSocket already has an active reader.", error.message) + end + + def test_each_enumerator_can_read_on_rubys_internal_fiber + socket = FakeSocket.new(text_delta("hello")) + + client.responses.connect(transport: FakeTransport.new(socket)) do |connection| + event = connection.each.next + assert_equal("hello", event.delta) + end + end + + def test_receive_returns_nil_after_local_close_without_reading_queued_data + socket = FakeSocket.new(text_delta("should-not-be-read")) + + client.responses.connect(transport: FakeTransport.new(socket)) do |connection| + connection.close + assert_nil(connection.receive) + end + end + + def test_receive_can_continue_after_malformed_json + socket = FakeSocket.new("{invalid", text_delta("still-readable")) + + client.responses.connect(transport: FakeTransport.new(socket)) do |connection| + assert_raises(OpenAI::Errors::ResponsesProtocolError) { connection.receive } + assert_equal("still-readable", connection.receive.delta) + end + end + + def test_foreign_owner_failure_does_not_release_the_active_read_lease + socket = FakeSocket.new(text_delta("hello"), text_delta("nested")) + + error = assert_raises(OpenAI::Errors::ResponsesConnectionError) do + client.responses.connect(transport: FakeTransport.new(socket)) do |connection| + connection.each do |_event| + thread = Thread.new do + assert_raises(OpenAI::Errors::ResponsesConnectionError) { connection.receive } + end + + thread.join + connection.receive + end + end + end + + assert_equal("Responses WebSocket already has an active reader.", error.message) + end + + def test_foreign_owner_close_and_abort_do_not_change_owner_state + %i[close abort].each do |operation| + socket = FakeSocket.new + + client.responses.connect(transport: FakeTransport.new(socket)) do |connection| + thread = Thread.new do + assert_raises(OpenAI::Errors::ResponsesConnectionError) do + operation == :close ? connection.close : connection.abort + end + end + + thread.join + + assert_nil(connection.response.create(model: "gpt-5.2")) + end + + assert_equal(1, socket.writes.size) + end + end + + def test_foreign_owner_send_does_not_write + socket = FakeSocket.new + + client.responses.connect(transport: FakeTransport.new(socket)) do |connection| + thread = Thread.new do + assert_raises(OpenAI::Errors::ResponsesConnectionError) do + connection.response.create(model: "gpt-5.2") + end + end + + thread.join + end + + assert_empty(socket.writes) + end + + def test_stream_ids_are_forwarded_without_client_policy + socket = FakeSocket.new + + client.responses.connect(transport: FakeTransport.new(socket)) do |connection| + connection.response.create(model: "gpt-5.2", stream_id: "bad id") + 33.times do |index| + assert_nil(connection.response.create(model: "gpt-5.2", stream_id: "lane_#{index}")) + end + end + + assert_equal(34, socket.writes.size) + end + + def test_response_fields_are_forwarded + socket = FakeSocket.new + + client.responses.connect(transport: FakeTransport.new(socket)) do |connection| + assert_nil(connection.response.create(model: "gpt-5.2", generate: false)) + assert_nil(connection.response.create(model: "gpt-5.2", background: true)) + end + + assert_equal(false, JSON.parse(socket.writes.fetch(0)).fetch("generate")) + assert_equal(true, JSON.parse(socket.writes.fetch(1)).fetch("background")) + end + + def test_explicit_base_url_is_validated_and_proxy_authorization_is_stripped + endpoint = +"wss://socket.example.test/custom/v2" + transport = FakeTransport.new(FakeSocket.new) + configured = client(default_headers: {"Proxy-Authorization" => "Basic configured-secret"}) + + configured + .responses + .connect( + websocket_base_url: endpoint, + request_options: {extra_headers: {"proxy-authorization" => "Basic request-secret"}}, + transport: transport + ) do |_connection| + endpoint.replace("wss://attacker.invalid/v1") + end + + assert_equal("wss://socket.example.test/custom/v2/responses", transport.open_args.fetch(:url).to_s) + refute(transport.open_args.fetch(:headers).key?("proxy-authorization")) + end + + def test_unsafe_base_urls_and_unsupported_request_options_fail_before_open + [ + "/socket", + "ftp://socket.example.test/v1", + "wss://user:secret@socket.example.test/v1", + "wss://socket.example.test/v1?tenant=one", + "wss://socket.example.test/v1#fragment" + ].each do |url| + transport = FakeTransport.new(FakeSocket.new) + assert_raises(ArgumentError) do + client.responses.connect(websocket_base_url: url, transport: transport) { |_connection| nil } + end + + assert_nil(transport.open_args) + end + + [{max_retries: 1}, {extra_query: {"token" => "secret"}}].each do |request_options| + transport = FakeTransport.new(FakeSocket.new) + assert_raises(ArgumentError) do + client.responses.connect(request_options: request_options, transport: transport) { |_connection| nil } + end + + assert_nil(transport.open_args) + end + end + + def test_provider_clients_are_rejected_before_open + azure = OpenAI::Client.new( + provider: OpenAI::Providers.azure( + endpoint: "https://resource.openai.azure.com", + api_key: "azure-key" + ) + ) + transport = FakeTransport.new(FakeSocket.new) + + error = assert_raises(OpenAI::Errors::Error) do + azure.responses.connect(transport: transport) { |_connection| nil } + end + + assert_equal("Responses WebSocket connections are not supported by providers.", error.message) + assert_nil(transport.open_args) + end + + def test_workload_identity_refreshes_once_after_pre_yield_401 + configured = workload_identity_client + transport = RejectOnceTransport.new + tokens = ["stale-token", "fresh-token"] + invalidations = 0 + + configured.workload_identity_auth.stub( + :get_token, + -> (deadline:) { + refute_nil(deadline) + tokens.shift + } + ) do + configured.workload_identity_auth.stub(:invalidate_token, -> { invalidations += 1 }) do + configured.responses.connect(transport: transport) { |_connection| nil } + end + end + + assert_equal(1, invalidations) + assert_equal(2, transport.attempts.length) + assert_equal("Bearer stale-token", transport.attempts.fetch(0).dig(:headers, "authorization")) + assert_equal("Bearer fresh-token", transport.attempts.fetch(1).dig(:headers, "authorization")) + end + + def test_workload_identity_timeout_stays_payload_free + configured = workload_identity_client(timeout: 0.01) + transport = FakeTransport.new(FakeSocket.new) + clock = [100.0, 100.02] + + error = OpenAI::Internal::Util.stub(:monotonic_secs, -> { clock.shift || 100.02 }) do + get_token = -> (deadline:) { + assert_equal(100.01, deadline) + "fresh-token" + } + + configured.workload_identity_auth.stub(:get_token, get_token) do + assert_raises(OpenAI::Errors::ResponsesConnectionError) do + configured.responses.connect(transport: transport) { |_connection| nil } + end + end + end + + assert_nil(error.cause) + refute_includes(error.full_message, "Timeout::Error") + assert_nil(transport.open_args) + end + + def test_inbound_stream_ids_are_forward_compatible + secret = "secret-stream-id" + socket = FakeSocket.new(JSON.generate(type: "response.future", stream_id: "bad #{secret}")) + + event = nil + client.responses.connect(transport: FakeTransport.new(socket)) { |connection| event = connection.receive } + + assert_equal("bad #{secret}", event.stream_id) + end + + def test_unknown_event_preserves_non_string_stream_id + socket = FakeSocket.new(JSON.generate(type: "response.future", stream_id: 123)) + event = nil + + client.responses.connect(transport: FakeTransport.new(socket)) { |connection| event = connection.receive } + + assert_equal(123, event.stream_id) + end + + def test_error_event_is_yielded_without_poisoning_connection + limit = JSON.generate( + type: "error", + error: { + type: "invalid_request_error", + code: "websocket_connection_limit_reached", + message: "Connection reached its limit.", + param: nil + } + ) + socket = FakeSocket.new(limit) + + client.responses.connect(transport: FakeTransport.new(socket)) do |connection| + event = connection.receive + assert_instance_of(OpenAI::Responses::ResponsesServerEvent::ResponseWsError, event) + assert_nil(connection.response.create(model: "gpt-5.2")) + end + end + + def test_default_transport_preserves_missing_dependency_guidance_without_cause + transport_class = Class.new(OpenAI::Responses::Transports::AsyncWebSocket) do + private def load_dependencies(url) + raise( + @error_factory.call( + url: url, + message: @dependency_message, + cause: LoadError.new("secret-load-path") + ) + ) + end + end + + transport = transport_class.new + + error = assert_raises(OpenAI::Errors::ResponsesConnectionError) do + transport.open(url: URI("wss://example.com/v1/responses"), headers: {}, timeout: 1) do + nil + end + end + + assert_equal( + "Responses WebSockets require the async-websocket gem. Add it to your Gemfile.", + error.message + ) + assert_nil(error.cause) + refute_includes(error.full_message, "secret-load-path") + end +end diff --git a/test/openai/responses_websocket/connection_test_support.rb b/test/openai/responses_websocket/connection_test_support.rb new file mode 100644 index 000000000..92f0ae6a8 --- /dev/null +++ b/test/openai/responses_websocket/connection_test_support.rb @@ -0,0 +1,120 @@ +# frozen_string_literal: true + +require_relative "../test_helper" + +module OpenAI::Test::ResponsesWebSocketConnectionTestSupport + class FakeSocket + attr_reader :writes, :close_args + + def initialize(*reads) + @reads = reads + @writes = [] + @closed = false + @aborted = false + end + + def read = @reads.shift + + def write(message) + @writes << message + nil + end + + def closed? = @closed + + def close(code: 1000, reason: "") + @closed = true + @close_args = {code: code, reason: reason} + nil + end + + def abort + @closed = true + @aborted = true + nil + end + + def aborted? = @aborted + end + + class FailingWriteSocket < FakeSocket + def write(_message) = raise IOError, "write failed with sensitive-body" + end + + class FakeTransport + attr_reader :open_args + + def initialize(socket) + @socket = socket + end + + def open(url:, headers:, timeout:, **options) + @open_args = {url: url, headers: headers, timeout: timeout, options: options} + yield(@socket) + end + end + + class RejectOnceTransport < FakeTransport + attr_reader :attempts + + def initialize + super(FakeSocket.new) + @attempts = [] + end + + def open(url:, headers:, timeout:, **options) + @attempts << {url: url, headers: headers, timeout: timeout, options: options} + if @attempts.one? + raise( + OpenAI::Errors::ResponsesConnectionError.new( + url: url, + message: "upgrade rejected", + http_status: 401 + ) + ) + end + + yield(@socket) + end + end + + private def client(**options) + OpenAI::Client.new( + api_key: "test-key", + base_url: "https://example.com/v1", + **options + ) + end + + private def workload_identity_client(timeout: 600) + provider = OpenAI::Auth::SubjectTokenProviders::K8sServiceAccountTokenProvider.new( + token_path: "/not-read-by-this-test" + ) + config = OpenAI::Auth::WorkloadIdentity.new( + identity_provider_id: "idp_123", + service_account_id: "sa_123", + provider: provider + ) + OpenAI::Client.new( + api_key: nil, + workload_identity: config, + organization: "org_123", + base_url: "https://example.com/v1", + timeout: timeout + ) + end + + private def text_delta(delta, stream_id: nil) + JSON.generate( + { + type: "response.output_text.delta", + sequence_number: 1, + item_id: "item_1", + output_index: 0, + content_index: 0, + delta: delta, + logprobs: [] + }.compact.merge(stream_id ? {stream_id: stream_id} : {}) + ) + end +end diff --git a/test/openai/responses_websocket/errors_test.rb b/test/openai/responses_websocket/errors_test.rb new file mode 100644 index 000000000..c38b50bab --- /dev/null +++ b/test/openai/responses_websocket/errors_test.rb @@ -0,0 +1,18 @@ +# frozen_string_literal: true + +require_relative "../test_helper" + +class OpenAI::Test::ResponsesWebSocketErrorsTest < Minitest::Test + def test_connection_errors_drop_credentials_and_query_without_mutating_url + url = URI("wss://user:secret@example.com/v1/responses?token=sensitive#fragment-secret") + original = url.to_s + + error = OpenAI::Errors::ResponsesConnectionError.new(url: url) + + assert_equal("wss://example.com/v1/responses", error.url.to_s) + assert_equal(original, url.to_s) + refute_includes(error.full_message, "sensitive") + refute_includes(error.url.to_s, "fragment-secret") + assert_nil(error.cause) + end +end diff --git a/test/openai/responses_websocket/sorbet_test.rb b/test/openai/responses_websocket/sorbet_test.rb new file mode 100644 index 000000000..0675cf87c --- /dev/null +++ b/test/openai/responses_websocket/sorbet_test.rb @@ -0,0 +1,66 @@ +# frozen_string_literal: true + +require "open3" +require "tempfile" + +require_relative "../test_helper" + +class OpenAI::Test::ResponsesWebSocketSorbetTest < Minitest::Test + def test_shipped_rbi_types_responses_websocket_helpers + source = <<~RUBY + # typed: strict + + connection = T.must(T.let(nil, T.nilable(OpenAI::Responses::Connection))) + + connection.response.create( + model: "gpt-5.2", + input: [OpenAI::Responses::EasyInputMessage.new(role: :user, content: "hello")], + stream_id: "turn_1" + ) + + event = connection.receive + T.assert_type!( + event, + T.nilable( + T.any( + OpenAI::Responses::ResponsesServerEvent::Variants, + OpenAI::Responses::UnknownServerEvent + ) + ) + ) + RUBY + + stdout, stderr, status = typecheck(source) + + assert_predicate(status, :success?, "#{stdout}\n#{stderr}") + end + + def test_shipped_rbi_rejects_wrong_known_response_create_type + source = <<~RUBY + # typed: strict + + connection = T.must(T.let(nil, T.nilable(OpenAI::Responses::Connection))) + connection.response.create(model: 123) + RUBY + + stdout, stderr, status = typecheck(source) + + refute_predicate(status, :success?, "#{stdout}\n#{stderr}") + assert_includes("#{stdout}\n#{stderr}", "Expected") + end + + private def typecheck(source) + root = File.expand_path("../../..", __dir__) + Tempfile.create(["responses-websocket-sorbet", ".rb"]) do |file| + file.write(source) + file.flush + Open3.capture3( + {"SRB_SKIP_GEM_RBIS" => "1"}, + "srb", + "typecheck", + file.path, + chdir: root + ) + end + end +end