diff --git a/core/http/endpoints/openresponses/responses.go b/core/http/endpoints/openresponses/responses.go index d86780fbd2c4..cb1de138e8b3 100644 --- a/core/http/endpoints/openresponses/responses.go +++ b/core/http/endpoints/openresponses/responses.go @@ -58,33 +58,17 @@ func ResponsesEndpoint(cl *config.ModelConfigLoader, ml *model.ModelLoader, eval shouldStore = false } - // Handle previous_response_id if provided - var previousResponse *schema.ORResponseResource + // Handle previous_response_id if provided. var messages []schema.Message if input.PreviousResponseID != "" { - stored, err := store.Get(input.PreviousResponseID) + previousMessages, err := resolvePreviousResponseMessages(store, input.PreviousResponseID, cfg) if err != nil { - return sendOpenResponsesError(c, 404, "not_found", fmt.Sprintf("previous response not found: %s", input.PreviousResponseID), "previous_response_id") - } - previousResponse = stored.Response - - // Also convert previous response input to messages - previousInputMessages, err := convertORInputToMessages(stored.Request.Input, cfg) - if err != nil { - return sendOpenResponsesError(c, 400, "invalid_request", fmt.Sprintf("failed to convert previous input: %v", err), "") - } - - // Convert previous response output items to messages - previousOutputMessages, err := convertOROutputItemsToMessages(previousResponse.Output) - if err != nil { - return sendOpenResponsesError(c, 400, "invalid_request", fmt.Sprintf("failed to convert previous response: %v", err), "") + if notFound, ok := err.(*previousResponseNotFoundError); ok { + return sendOpenResponsesError(c, 404, "not_found", notFound.Error(), "previous_response_id") + } + return sendOpenResponsesError(c, 400, "invalid_request", err.Error(), "") } - - // Concatenate: previous_input + previous_output + new_input - // Start with previous input messages - messages = previousInputMessages - // Add previous output as assistant messages - messages = append(messages, previousOutputMessages...) + messages = previousMessages } // Convert Open Responses input to internal Messages @@ -266,7 +250,7 @@ func ResponsesEndpoint(cl *config.ModelConfigLoader, ml *model.ModelLoader, eval if input.Stream { // Background streaming processing (buffer events) - finalResponse, bgErr = handleBackgroundStream(bgCtx, store, responseID, createdAt, input, cfg, ml, cl, appConfig, predInput, openAIReq, funcs, shouldUseFn, mcpExecutor, evaluator) + finalResponse, bgErr = handleBackgroundStream(bgCtx, store, responseID, createdAt, input, cfg, ml, cl, appConfig, predInput, openAIReq, funcs, shouldUseFn, true, mcpExecutor, evaluator) } else { // Background non-streaming processing finalResponse, bgErr = handleBackgroundNonStream(bgCtx, store, responseID, createdAt, input, cfg, ml, cl, appConfig, predInput, openAIReq, funcs, shouldUseFn, mcpExecutor, evaluator) @@ -515,6 +499,76 @@ func extractReasoningContentFromORItem(item *schema.ORItemField) string { return "" } +type previousResponseNotFoundError struct { + ResponseID string +} + +func (e *previousResponseNotFoundError) Error() string { + return fmt.Sprintf("previous response not found: %s", e.ResponseID) +} + +// resolvePreviousResponseMessages reconstructs the complete stored conversation +// ending at responseID. Requests are stored as incremental deltas, so replaying +// only the immediately previous request loses older turns after the first chain. +func resolvePreviousResponseMessages(store *ResponseStore, responseID string, cfg *config.ModelConfig) ([]schema.Message, error) { + return resolvePreviousResponseMessagesFromStores([]*ResponseStore{store}, responseID, cfg) +} + +// resolvePreviousResponseMessagesFromStores resolves each hop against the stores +// in priority order. WebSocket mode uses this to prefer connection-local +// store=false responses while still allowing references to globally stored ones. +func resolvePreviousResponseMessagesFromStores(stores []*ResponseStore, responseID string, cfg *config.ModelConfig) ([]schema.Message, error) { + type chainEntry struct { + id string + stored *StoredResponse + } + + var chain []chainEntry + seen := make(map[string]struct{}) + for currentID := responseID; currentID != ""; { + if _, exists := seen[currentID]; exists { + return nil, fmt.Errorf("previous_response_id cycle detected at %s", currentID) + } + seen[currentID] = struct{}{} + + var stored *StoredResponse + for _, store := range stores { + if store == nil { + continue + } + candidate, err := store.Get(currentID) + if err == nil { + stored = candidate + break + } + } + if stored == nil { + return nil, &previousResponseNotFoundError{ResponseID: currentID} + } + if stored.Request == nil || stored.Response == nil { + return nil, fmt.Errorf("stored previous response %s is incomplete", currentID) + } + chain = append(chain, chainEntry{id: currentID, stored: stored}) + currentID = stored.Request.PreviousResponseID + } + + var messages []schema.Message + for i := len(chain) - 1; i >= 0; i-- { + entry := chain[i] + inputMessages, err := convertORInputToMessages(entry.stored.Request.Input, cfg) + if err != nil { + return nil, fmt.Errorf("failed to convert previous input for %s: %w", entry.id, err) + } + outputMessages, err := convertOROutputItemsToMessages(entry.stored.Response.Output) + if err != nil { + return nil, fmt.Errorf("failed to convert previous response %s: %w", entry.id, err) + } + messages = append(messages, inputMessages...) + messages = append(messages, outputMessages...) + } + return messages, nil +} + // convertOROutputItemsToMessages converts Open Responses output items to internal Messages. // Contiguous assistant items (message, reasoning, function_call) are merged into a single message. func convertOROutputItemsToMessages(outputItems []schema.ORItemField) ([]schema.Message, error) { @@ -1013,7 +1067,7 @@ func handleBackgroundNonStream(ctx context.Context, store *ResponseStore, respon } // handleBackgroundStream handles background streaming responses with event buffering -func handleBackgroundStream(ctx context.Context, store *ResponseStore, responseID string, createdAt int64, input *schema.OpenResponsesRequest, cfg *config.ModelConfig, ml *model.ModelLoader, cl *config.ModelConfigLoader, appConfig *config.ApplicationConfig, predInput string, openAIReq *schema.OpenAIRequest, funcs functions.Functions, shouldUseFn bool, mcpExecutor mcpTools.ToolExecutor, evaluator *templates.Evaluator) (*schema.ORResponseResource, error) { +func handleBackgroundStream(ctx context.Context, store *ResponseStore, responseID string, createdAt int64, input *schema.OpenResponsesRequest, cfg *config.ModelConfig, ml *model.ModelLoader, cl *config.ModelConfigLoader, appConfig *config.ApplicationConfig, predInput string, openAIReq *schema.OpenAIRequest, funcs functions.Functions, shouldUseFn bool, shouldStore bool, mcpExecutor mcpTools.ToolExecutor, evaluator *templates.Evaluator) (*schema.ORResponseResource, error) { // Populate openAIReq fields for ComputeChoices openAIReq.Tools = convertORToolsToOpenAIFormat(input.Tools) openAIReq.ToolsChoice = input.ToolChoice @@ -1026,7 +1080,7 @@ func handleBackgroundStream(ctx context.Context, store *ResponseStore, responseI sequenceNumber := 0 // Emit response.created - responseCreated := buildORResponse(responseID, createdAt, nil, schema.ORStatusInProgress, input, []schema.ORItemField{}, nil, true) + responseCreated := buildORResponse(responseID, createdAt, nil, schema.ORStatusInProgress, input, []schema.ORItemField{}, nil, shouldStore) bufferEvent(store, responseID, &schema.ORStreamEvent{ Type: "response.created", SequenceNumber: sequenceNumber, @@ -1298,7 +1352,7 @@ func handleBackgroundStream(ctx context.Context, store *ResponseStore, responseI InputTokens: lastTokenUsage.Prompt, OutputTokens: lastTokenUsage.Completion, TotalTokens: lastTokenUsage.Prompt + lastTokenUsage.Completion, - }, true) + }, shouldStore) // Emit response.completed bufferEvent(store, responseID, &schema.ORStreamEvent{ @@ -2955,12 +3009,15 @@ func sendOpenResponsesError(c echo.Context, statusCode int, errorType, message, return c.JSON(statusCode, errorResp) } -// convertORToolsToOpenAIFormat converts Open Responses tools to OpenAI format for the backend -// Open Responses format: { type, name, description, parameters } -// OpenAI format: { type, function: { name, description, parameters } } +// convertORToolsToOpenAIFormat converts only tools that have an equivalent in +// the OpenAI-compatible function-tool representation. Native Responses tools +// such as web_search and namespace must not be rewritten as functions. func convertORToolsToOpenAIFormat(orTools []schema.ORFunctionTool) []functions.Tool { result := make([]functions.Tool, 0, len(orTools)) for _, t := range orTools { + if t.Type != "function" { + continue + } result = append(result, functions.Tool{ Type: "function", Function: functions.Function{ diff --git a/core/http/endpoints/openresponses/responses_convert_test.go b/core/http/endpoints/openresponses/responses_convert_test.go index 9dfe861bb3d3..0947a80b9ff2 100644 --- a/core/http/endpoints/openresponses/responses_convert_test.go +++ b/core/http/endpoints/openresponses/responses_convert_test.go @@ -2,6 +2,7 @@ package openresponses import ( "github.com/mudler/LocalAI/core/config" + "github.com/mudler/LocalAI/core/schema" . "github.com/onsi/ginkgo/v2" . "github.com/onsi/gomega" @@ -60,3 +61,78 @@ var _ = Describe("convertORInputToMessages", func() { Expect(msgs).To(BeEmpty()) }) }) + +var _ = Describe("convertORToolsToOpenAIFormat", func() { + It("only converts Responses function tools", func() { + converted := convertORToolsToOpenAIFormat([]schema.ORFunctionTool{ + {Type: "function", Name: "example_function", Parameters: map[string]any{"type": "object"}}, + {Type: "web_search"}, + {Type: "namespace", Name: "multi_agent_v1"}, + }) + + Expect(converted).To(HaveLen(1)) + Expect(converted[0].Type).To(Equal("function")) + Expect(converted[0].Function.Name).To(Equal("example_function")) + Expect(converted[0].Function.Parameters).To(Equal(map[string]any{"type": "object"})) + }) +}) + +var _ = Describe("resolvePreviousResponseMessages", func() { + It("replays multi-hop response history from oldest to newest", func() { + store := NewResponseStore(0) + cfg := &config.ModelConfig{} + message := func(role, text string) schema.ORItemField { + return schema.ORItemField{ + Type: "message", + Role: role, + Content: []schema.ORContentPart{{Type: "output_text", Text: text}}, + } + } + + store.Store("resp_0", &schema.OpenResponsesRequest{Input: "base"}, &schema.ORResponseResource{ + ID: "resp_0", Output: []schema.ORItemField{message("assistant", "answer-0")}, + }) + store.Store("resp_1", &schema.OpenResponsesRequest{PreviousResponseID: "resp_0", Input: "question-1"}, &schema.ORResponseResource{ + ID: "resp_1", Output: []schema.ORItemField{message("assistant", "answer-1")}, + }) + store.Store("resp_2", &schema.OpenResponsesRequest{PreviousResponseID: "resp_1", Input: "question-2"}, &schema.ORResponseResource{ + ID: "resp_2", Output: []schema.ORItemField{message("assistant", "answer-2")}, + }) + + msgs, err := resolvePreviousResponseMessages(store, "resp_2", cfg) + Expect(err).NotTo(HaveOccurred()) + Expect(msgs).To(HaveLen(6)) + Expect([]string{ + msgs[0].StringContent, msgs[1].StringContent, + msgs[2].StringContent, msgs[3].StringContent, + msgs[4].StringContent, msgs[5].StringContent, + }).To(Equal([]string{"base", "answer-0", "question-1", "answer-1", "question-2", "answer-2"})) + }) + + It("resolves a chain across connection-local and global stores", func() { + connectionStore := NewResponseStore(0) + globalStore := NewResponseStore(0) + cfg := &config.ModelConfig{} + message := func(text string) schema.ORItemField { + return schema.ORItemField{ + Type: "message", + Role: "assistant", + Content: []schema.ORContentPart{{Type: "output_text", Text: text}}, + } + } + + globalStore.Store("resp_global", &schema.OpenResponsesRequest{Input: "base"}, &schema.ORResponseResource{ + ID: "resp_global", Output: []schema.ORItemField{message("answer-0")}, + }) + connectionStore.Store("resp_local", &schema.OpenResponsesRequest{PreviousResponseID: "resp_global", Input: "question-1"}, &schema.ORResponseResource{ + ID: "resp_local", Output: []schema.ORItemField{message("answer-1")}, + }) + + msgs, err := resolvePreviousResponseMessagesFromStores([]*ResponseStore{connectionStore, globalStore}, "resp_local", cfg) + Expect(err).NotTo(HaveOccurred()) + Expect(msgs).To(HaveLen(4)) + Expect([]string{msgs[0].StringContent, msgs[1].StringContent, msgs[2].StringContent, msgs[3].StringContent}).To( + Equal([]string{"base", "answer-0", "question-1", "answer-1"}), + ) + }) +}) diff --git a/core/http/endpoints/openresponses/websocket.go b/core/http/endpoints/openresponses/websocket.go index ffff7b0445c1..370fe98a97c8 100644 --- a/core/http/endpoints/openresponses/websocket.go +++ b/core/http/endpoints/openresponses/websocket.go @@ -44,6 +44,16 @@ func (lc *lockedConn) writeJSON(v any) error { return lc.Conn.WriteJSON(v) } +// writeTerminalJSON atomically hands the connection to the next response: the +// in-flight guard is released while the write lock is held, so a newly accepted +// response cannot write an event ahead of this terminal event. +func (lc *lockedConn) writeTerminalJSON(v any, release func()) error { + lc.Lock() + defer lc.Unlock() + release() + return lc.Conn.WriteJSON(v) +} + // WebSocketEndpoint handles WebSocket mode for the Responses API. // Clients connect via ws://:/v1/responses and send response.create messages. // Events are streamed back over the WebSocket connection instead of SSE. @@ -82,6 +92,16 @@ func WebSocketEndpoint(application *application.Application) echo.HandlerFunc { // handleWebSocketConnection runs the read loop for a single WebSocket connection. func handleWebSocketConnection(connCtx context.Context, conn *lockedConn, cl *config.ModelConfigLoader, ml *model.ModelLoader, evaluator *templates.Evaluator, appConfig *config.ApplicationConfig) { + // Responses created with store=false remain available only to this WebSocket + // connection, matching the Responses WebSocket continuation contract without + // leaking zero-data-retention state into the process-wide response store. + connectionStore := NewResponseStore(0) + defer func() { + if err := connectionStore.Close(); err != nil { + xlog.Warn("WebSocket Responses: failed to close connection-local response store", "error", err) + } + }() + // Track in-flight response to enforce one-at-a-time var inflight sync.Mutex @@ -130,8 +150,12 @@ func handleWebSocketConnection(connCtx context.Context, conn *lockedConn, cl *co } go func() { - defer inflight.Unlock() - handleWSResponseCreate(connCtx, conn, &wsMsg.OpenResponsesRequest, cl, ml, evaluator, appConfig) + var releaseOnce sync.Once + release := func() { + releaseOnce.Do(inflight.Unlock) + } + defer release() + handleWSResponseCreate(connCtx, conn, connectionStore, release, wsMsg.Generate, &wsMsg.OpenResponsesRequest, cl, ml, evaluator, appConfig) }() } } @@ -140,12 +164,18 @@ func handleWebSocketConnection(connCtx context.Context, conn *lockedConn, cl *co // It reuses the existing background stream infrastructure: the request is processed via // handleBackgroundStream which buffers events into the store, and a forwarder goroutine // reads those events and sends them over the WebSocket. -func handleWSResponseCreate(connCtx context.Context, conn *lockedConn, input *schema.OpenResponsesRequest, cl *config.ModelConfigLoader, ml *model.ModelLoader, evaluator *templates.Evaluator, appConfig *config.ApplicationConfig) { +func handleWSResponseCreate(connCtx context.Context, conn *lockedConn, connectionStore *ResponseStore, release func(), generate *bool, input *schema.OpenResponsesRequest, cl *config.ModelConfigLoader, ml *model.ModelLoader, evaluator *templates.Evaluator, appConfig *config.ApplicationConfig) { createdAt := time.Now().Unix() responseID := fmt.Sprintf("resp_%s", uuid.New().String()) + fail := func(errType, message, param string) { + sendWSErrorAndRelease(conn, release, errType, message, param) + } + failEvent := func(code, message, param string) { + sendWSErrorEventAndRelease(conn, release, code, message, param) + } if input.Model == "" { - sendWSError(conn, "invalid_request", "model is required", "model") + fail("invalid_request", "model is required", "model") return } @@ -153,7 +183,7 @@ func handleWSResponseCreate(connCtx context.Context, conn *lockedConn, input *sc cfg, err := cl.LoadModelConfigFileByNameDefaultOptions(input.Model, appConfig) if err != nil { xlog.Warn("WebSocket Responses: model config not found", "model", input.Model, "error", err) - sendWSError(conn, "invalid_request", fmt.Sprintf("model not found: %s", input.Model), "model") + fail("invalid_request", fmt.Sprintf("model not found: %s", input.Model), "model") return } if cfg.Model == "" { @@ -162,7 +192,7 @@ func handleWSResponseCreate(connCtx context.Context, conn *lockedConn, input *sc // Merge request params into config (same as mergeOpenResponsesRequestAndModelConfig) if err := middleware.MergeOpenResponsesConfig(cfg, input); err != nil { - sendWSError(conn, "invalid_request", fmt.Sprintf("invalid configuration: %v", err), "") + fail("invalid_request", fmt.Sprintf("invalid configuration: %v", err), "") return } @@ -173,9 +203,9 @@ func handleWSResponseCreate(connCtx context.Context, conn *lockedConn, input *sc input.Context = reqCtx input.Cancel = reqCancel - store := GetGlobalStore() + globalStore := GetGlobalStore() if appConfig.OpenResponsesStoreTTL > 0 { - store.SetTTL(appConfig.OpenResponsesStoreTTL) + globalStore.SetTTL(appConfig.OpenResponsesStoreTTL) } shouldStore := true @@ -183,36 +213,64 @@ func handleWSResponseCreate(connCtx context.Context, conn *lockedConn, input *sc shouldStore = false } - // Handle previous_response_id - var messages []schema.Message - if input.PreviousResponseID != "" { - stored, err := store.Get(input.PreviousResponseID) - if err != nil { - sendWSErrorEvent(conn, "previous_response_not_found", - fmt.Sprintf("previous response not found: %s", input.PreviousResponseID), "previous_response_id") - return - } + store := globalStore + if !shouldStore { + store = connectionStore + } - previousInputMessages, err := convertORInputToMessages(stored.Request.Input, cfg) - if err != nil { - sendWSError(conn, "invalid_request", fmt.Sprintf("failed to convert previous input: %v", err), "") + // Codex uses generate=false to prewarm the Responses WebSocket with the + // exact request it may send next. Persist the request and return a terminal + // response ID, but do not build a prompt or invoke the model backend. + if generate != nil && !*generate { + responseCreated := buildORResponse(responseID, createdAt, nil, schema.ORStatusInProgress, input, []schema.ORItemField{}, nil, shouldStore) + store.StoreBackground(responseID, input, responseCreated, reqCancel, true) + bufferEvent(store, responseID, &schema.ORStreamEvent{ + Type: "response.created", + SequenceNumber: 0, + Response: responseCreated, + }) + + now := time.Now().Unix() + responseCompleted := buildORResponse(responseID, createdAt, &now, schema.ORStatusCompleted, input, []schema.ORItemField{}, nil, shouldStore) + if err := store.UpdateResponse(responseID, responseCompleted); err != nil { + fail("server_error", fmt.Sprintf("failed to complete prewarm response: %v", err), "") + if !shouldStore { + store.Delete(responseID) + } return } + bufferEvent(store, responseID, &schema.ORStreamEvent{ + Type: "response.completed", + SequenceNumber: 1, + Response: responseCompleted, + }) + + processDone := make(chan struct{}) + close(processDone) + forwardEvents(reqCtx, conn, store, responseID, processDone, release) + return + } - previousOutputMessages, err := convertOROutputItemsToMessages(stored.Response.Output) + // Handle previous_response_id. WebSocket continuations may refer to either + // connection-local store=false responses or process-wide stored responses. + var messages []schema.Message + if input.PreviousResponseID != "" { + previousMessages, err := resolvePreviousResponseMessagesFromStores([]*ResponseStore{connectionStore, globalStore}, input.PreviousResponseID, cfg) if err != nil { - sendWSError(conn, "invalid_request", fmt.Sprintf("failed to convert previous response: %v", err), "") + if notFound, ok := err.(*previousResponseNotFoundError); ok { + failEvent("previous_response_not_found", notFound.Error(), "previous_response_id") + return + } + fail("invalid_request", err.Error(), "") return } - - messages = previousInputMessages - messages = append(messages, previousOutputMessages...) + messages = previousMessages } // Convert current input to messages newMessages, err := convertORInputToMessages(input.Input, cfg) if err != nil { - sendWSError(conn, "invalid_request", fmt.Sprintf("failed to parse input: %v", err), "") + fail("invalid_request", fmt.Sprintf("failed to parse input: %v", err), "") return } messages = append(messages, newMessages...) @@ -308,7 +366,7 @@ func handleWSResponseCreate(connCtx context.Context, conn *lockedConn, input *sc defer close(processDone) store.UpdateStatus(responseID, schema.ORStatusInProgress, nil) - finalResponse, bgErr := handleBackgroundStream(reqCtx, store, responseID, createdAt, input, cfg, ml, cl, appConfig, predInput, openAIReq, funcs, shouldUseFn, nil, nil) + finalResponse, bgErr := handleBackgroundStream(reqCtx, store, responseID, createdAt, input, cfg, ml, cl, appConfig, predInput, openAIReq, funcs, shouldUseFn, shouldStore, nil, nil) if bgErr != nil { xlog.Error("WebSocket Responses: processing failed", "response_id", responseID, "error", bgErr) now := time.Now().Unix() @@ -332,21 +390,37 @@ func handleWSResponseCreate(connCtx context.Context, conn *lockedConn, input *sc }() // Forward events from the store to the WebSocket connection - forwardEvents(reqCtx, conn, store, responseID, processDone, shouldStore) + forwardEvents(reqCtx, conn, store, responseID, processDone, release) } // forwardEvents subscribes to events for a response and sends them over the WebSocket. // This mirrors handleStreamResume but writes JSON to WebSocket instead of SSE. -func forwardEvents(ctx context.Context, conn *lockedConn, store *ResponseStore, responseID string, done <-chan struct{}, shouldStore bool) { +func forwardEvents(ctx context.Context, conn *lockedConn, store *ResponseStore, responseID string, done <-chan struct{}, release func()) { eventsChan, err := store.GetEventsChan(responseID) if err != nil { return } - lastSeq := -1 + writeEvent := func(parsed *schema.ORStreamEvent) (terminal bool, err error) { + switch parsed.Type { + case "response.completed", "response.failed", "error": + // A terminal event is the protocol boundary for accepting the next + // response.create. Wait until processing has fully stopped, then + // release the in-flight guard before making that event visible. + select { + case <-ctx.Done(): + return true, ctx.Err() + case <-done: + } + return true, conn.writeTerminalJSON(parsed, release) + default: + return false, conn.writeJSON(parsed) + } + } + lastSeq := -1 for { - // Drain all available events + // Drain all available events. events, err := store.GetEventsAfter(responseID, lastSeq) if err != nil { return @@ -356,16 +430,16 @@ func forwardEvents(ctx context.Context, conn *lockedConn, store *ResponseStore, if err := json.Unmarshal(event.Data, &parsed); err != nil { continue } - if err := conn.writeJSON(&parsed); err != nil { + terminal, err := writeEvent(&parsed) + if err != nil || terminal { return } lastSeq = event.SequenceNumber } - // Check if processing is done and all events have been sent + // Check if processing is done and all events have been sent. select { case <-done: - // Drain any final events finalEvents, err := store.GetEventsAfter(responseID, lastSeq) if err == nil { for _, event := range finalEvents { @@ -373,27 +447,24 @@ func forwardEvents(ctx context.Context, conn *lockedConn, store *ResponseStore, if err := json.Unmarshal(event.Data, &parsed); err != nil { continue } - if err := conn.writeJSON(&parsed); err != nil { + terminal, err := writeEvent(&parsed) + if err != nil || terminal { return } } } - // Clean up non-stored responses from the cache - if !shouldStore { - store.Delete(responseID) - } return default: } - // Wait for new events, completion, or context cancellation + // Wait for new events, completion, or context cancellation. select { case <-ctx.Done(): return case <-done: - // Will drain in next iteration + // Will drain in next iteration. case <-eventsChan: - // New events available + // New events available. } } } @@ -410,6 +481,20 @@ func sendWSError(conn *lockedConn, errType, message, param string) { conn.writeJSON(&event) } +func sendWSErrorAndRelease(conn *lockedConn, release func(), errType, message, param string) { + event := schema.ORStreamEvent{ + Type: "error", + Error: &schema.ORErrorPayload{ + Type: errType, + Message: message, + Param: param, + }, + } + if err := conn.writeTerminalJSON(&event, release); err != nil { + xlog.Debug("WebSocket Responses: failed to write terminal error", "error", err) + } +} + func sendWSErrorEvent(conn *lockedConn, code, message, param string) { event := schema.ORStreamEvent{ Type: "error", @@ -422,3 +507,18 @@ func sendWSErrorEvent(conn *lockedConn, code, message, param string) { } conn.writeJSON(&event) } + +func sendWSErrorEventAndRelease(conn *lockedConn, release func(), code, message, param string) { + event := schema.ORStreamEvent{ + Type: "error", + Error: &schema.ORErrorPayload{ + Type: "invalid_request_error", + Code: code, + Message: message, + Param: param, + }, + } + if err := conn.writeTerminalJSON(&event, release); err != nil { + xlog.Debug("WebSocket Responses: failed to write terminal error", "error", err) + } +} diff --git a/core/schema/openresponses.go b/core/schema/openresponses.go index 98c57857bc3e..4cc0e05b3575 100644 --- a/core/schema/openresponses.go +++ b/core/schema/openresponses.go @@ -15,10 +15,12 @@ const ( ) // ORWebSocketMessage is the envelope for WebSocket mode messages. -// The client sends {"type":"response.create", ...} where the remaining fields -// map to OpenResponsesRequest. "type" is the only additional field. +// The client sends {"type":"response.create", ...} where most remaining fields +// map to OpenResponsesRequest. generate is a WebSocket-only control used by +// clients such as Codex to prewarm a request without running inference. type ORWebSocketMessage struct { - Type string `json:"type"` + Type string `json:"type"` + Generate *bool `json:"generate,omitempty"` OpenResponsesRequest } @@ -68,9 +70,11 @@ func (r *OpenResponsesRequest) ModelName(s *string) string { return r.Model } -// ORFunctionTool represents a function tool definition +// ORFunctionTool stores the function-shaped fields LocalAI consumes from a +// Responses tool entry. Type can be a native Responses tool kind; unsupported +// native fields are ignored and must not be reinterpreted as function fields. type ORFunctionTool struct { - Type string `json:"type"` // always "function" + Type string `json:"type"` Name string `json:"name"` Description string `json:"description,omitempty"` Parameters map[string]any `json:"parameters,omitempty"` diff --git a/tests/e2e/e2e_websocket_responses_test.go b/tests/e2e/e2e_websocket_responses_test.go index a25abf32d195..80003fb2be8d 100644 --- a/tests/e2e/e2e_websocket_responses_test.go +++ b/tests/e2e/e2e_websocket_responses_test.go @@ -33,6 +33,7 @@ type wsResponseBody struct { ID string `json:"id"` Status string `json:"status"` Model string `json:"model"` + Store bool `json:"store"` Output []struct { Type string `json:"type"` ID string `json:"id"` @@ -66,7 +67,7 @@ func readAllEvents(conn *websocket.Conn) []wsEvent { break } events = append(events, ev) - if ev.Type == "response.completed" || ev.Type == "response.failed" { + if ev.Type == "response.completed" || ev.Type == "response.failed" || ev.Type == "error" { break } } @@ -127,6 +128,73 @@ var _ = Describe("WebSocket Responses API E2E Tests", Label("WebSocket"), func() }) }) + Context("Codex WebSocket prewarm", func() { + It("does not generate for generate:false and reuses the response ID", func() { + conn, err := dialWS() + Expect(err).ToNot(HaveOccurred()) + defer func() { _ = conn.Close() }() + + warmup := map[string]any{ + "type": "response.create", + "model": "mock-model", + "store": false, + "generate": false, + "input": []map[string]any{{ + "type": "message", + "role": "user", + "content": []map[string]any{{"type": "input_text", "text": "Hello from Codex"}}, + }}, + } + Expect(conn.WriteJSON(warmup)).To(Succeed()) + + warmupEvents := readAllEvents(conn) + Expect(warmupEvents).To(HaveLen(2)) + Expect([]string{warmupEvents[0].Type, warmupEvents[1].Type}).To(Equal([]string{ + "response.created", + "response.completed", + })) + + var warmupResp wsResponseBody + lastWarmup := warmupEvents[len(warmupEvents)-1] + Expect(lastWarmup.Type).To(Equal("response.completed")) + Expect(json.Unmarshal(lastWarmup.Response, &warmupResp)).To(Succeed()) + Expect(warmupResp.ID).ToNot(BeEmpty()) + Expect(warmupResp.Store).To(BeFalse()) + Expect(warmupResp.Output).To(BeEmpty()) + + followUp := map[string]any{ + "type": "response.create", + "model": "mock-model", + "store": false, + "previous_response_id": warmupResp.ID, + "input": []any{}, + } + + // store:false response IDs are scoped to the WebSocket connection. + otherConn, err := dialWS() + Expect(err).ToNot(HaveOccurred()) + Expect(otherConn.WriteJSON(followUp)).To(Succeed()) + otherEvent, err := readEvent(otherConn) + Expect(err).ToNot(HaveOccurred()) + Expect(otherEvent.Type).To(Equal("error")) + Expect(otherEvent.Error).ToNot(BeNil()) + Expect(otherEvent.Error.Code).To(Equal("previous_response_not_found")) + Expect(otherConn.Close()).To(Succeed()) + + Expect(conn.WriteJSON(followUp)).To(Succeed()) + + followUpEvents := readAllEvents(conn) + Expect(followUpEvents).ToNot(BeEmpty()) + lastFollowUp := followUpEvents[len(followUpEvents)-1] + Expect(lastFollowUp.Type).To(Equal("response.completed"), "unexpected terminal event: %#v", lastFollowUp.Error) + var followUpResp wsResponseBody + Expect(json.Unmarshal(lastFollowUp.Response, &followUpResp)).To(Succeed()) + Expect(followUpResp.ID).ToNot(Equal(warmupResp.ID)) + Expect(followUpResp.Store).To(BeFalse()) + Expect(followUpResp.Output).ToNot(BeEmpty()) + }) + }) + Context("Continuation with previous_response_id", func() { It("chains responses using previous_response_id", func() { conn, err := dialWS()