fix(openresponses): support Codex WebSocket warm-up (#11608)

* fix(openresponses): support Codex WebSocket warmup

Assisted-by: ChatGPT:GPT-5.6-Sol golangci-lint
Signed-off-by: Abdullah Mansour <abdullahmansour.marketing@gmail.com>

* docs(openresponses): document WebSocket responses

Assisted-by: Codex:GPT-5.6-Sol gh Docker golangci-lint
Signed-off-by: Abdullah Mansour <abdullahmansour.marketing@gmail.com>

* fix(openresponses): harden WebSocket response lifecycle

Ensure response ownership, continuation storage, error sequencing, and connection-local resource limits remain correct across HTTP and WebSocket transports.

Assisted-by: Codex:GPT-5.6-Sol [gh] [Docker] [golangci-lint]
Signed-off-by: Abdullah Mansour <abdullahmansour.marketing@gmail.com>

---------

Signed-off-by: Abdullah Mansour <abdullahmansour.marketing@gmail.com>
This commit is contained in:
Abdullah Mansour | عبد الله منصور authored and GitHub committed 2026-09-11 16:54:42 +02:00
1 parent c0993e580a
commit acc7284526
9 files changed
+1024 -107

No files matched your search

+119 -37
View File
@@ -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, ownerFromContext(c))
if err != nil {
return sendOpenResponsesError(c, 404, "not_found", fmt.Sprintf("previous response not found: %s", input.PreviousResponseID), "previous_response_id")
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(), "")
}
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), "")
}
// 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
@@ -251,8 +235,7 @@ func ResponsesEndpoint(cl *config.ModelConfigLoader, ml *model.ModelLoader, eval
// Store the background response and stamp its owner before the ID
// is returned to the client, so later GET/cancel/resume can verify
// the caller owns it.
store.StoreBackground(responseID, input, queuedResponse, bgCancel, input.Stream)
store.SetOwner(responseID, ownerFromContext(c))
store.StoreBackgroundOwned(responseID, input, queuedResponse, bgCancel, input.Stream, ownerFromContext(c))
// Start background processing goroutine
go func() {
@@ -266,7 +249,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 +498,104 @@ 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, callerID string) ([]schema.Message, error) {
messages, _, err := resolvePreviousResponseMessagesFromSources(
[]previousResponseStoreSource{{store: store}},
responseID,
cfg,
callerID,
)
return messages, err
}
type previousResponseStoreSource struct {
store *ResponseStore
connectionLocal bool
}
// resolvePreviousResponseMessagesFromSources resolves each hop against the
// stores in priority order and reports whether the resulting chain contains
// connection-local state. Every hop is owner-checked so a known response ID
// cannot be used to replay another caller's conversation.
func resolvePreviousResponseMessagesFromSources(sources []previousResponseStoreSource, responseID string, cfg *config.ModelConfig, callerID string) ([]schema.Message, bool, error) {
type chainEntry struct {
id string
request schema.OpenResponsesRequest
response schema.ORResponseResource
}
var chain []chainEntry
usedConnectionLocal := false
seen := make(map[string]struct{})
for currentID := responseID; currentID != ""; {
if _, exists := seen[currentID]; exists {
return nil, false, fmt.Errorf("previous_response_id cycle detected at %s", currentID)
}
seen[currentID] = struct{}{}
var stored *StoredResponse
var connectionLocal bool
for _, source := range sources {
if source.store == nil {
continue
}
candidate, err := source.store.Get(currentID)
if err == nil {
if !accessAllowed(candidate, callerID) {
return nil, false, &previousResponseNotFoundError{ResponseID: currentID}
}
stored = candidate
connectionLocal = source.connectionLocal
break
}
}
if stored == nil {
return nil, false, &previousResponseNotFoundError{ResponseID: currentID}
}
stored.mu.RLock()
if stored.Request == nil || stored.Response == nil {
stored.mu.RUnlock()
return nil, false, fmt.Errorf("stored previous response %s is incomplete", currentID)
}
request := *stored.Request
response := *stored.Response
response.Output = append([]schema.ORItemField(nil), stored.Response.Output...)
stored.mu.RUnlock()
chain = append(chain, chainEntry{id: currentID, request: request, response: response})
usedConnectionLocal = usedConnectionLocal || connectionLocal
currentID = request.PreviousResponseID
}
var messages []schema.Message
for i := len(chain) - 1; i >= 0; i-- {
entry := chain[i]
inputMessages, err := convertORInputToMessages(entry.request.Input, cfg)
if err != nil {
return nil, false, fmt.Errorf("failed to convert previous input for %s: %w", entry.id, err)
}
outputMessages, err := convertOROutputItemsToMessages(entry.response.Output)
if err != nil {
return nil, false, fmt.Errorf("failed to convert previous response %s: %w", entry.id, err)
}
messages = append(messages, inputMessages...)
messages = append(messages, outputMessages...)
}
return messages, usedConnectionLocal, 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 +1094,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 +1107,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 +1379,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{
@@ -1591,8 +1672,7 @@ func handleOpenResponsesNonStream(c echo.Context, responseID string, createdAt i
// Store response for future reference (if enabled)
if shouldStore {
store := GetGlobalStore()
store.Store(responseID, input, response)
store.SetOwner(responseID, ownerFromContext(c))
store.StoreOwned(responseID, input, response, ownerFromContext(c))
}
return c.JSON(200, response)
@@ -2327,8 +2407,7 @@ func handleOpenResponsesStream(c echo.Context, responseID string, createdAt int6
// Store response for future reference (if enabled)
if shouldStore {
store := GetGlobalStore()
store.Store(responseID, input, responseCompleted)
store.SetOwner(responseID, ownerFromContext(c))
store.StoreOwned(responseID, input, responseCompleted, ownerFromContext(c))
}
// Send [DONE]
@@ -2683,7 +2762,7 @@ func handleOpenResponsesStream(c echo.Context, responseID string, createdAt int6
// Store response for future reference (if enabled)
if shouldStore {
store := GetGlobalStore()
store.Store(responseID, input, responseCompleted)
store.StoreOwned(responseID, input, responseCompleted, ownerFromContext(c))
}
// Send [DONE]
@@ -2955,12 +3034,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{
@@ -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,86 @@ 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 := resolvePreviousResponseMessagesFromSources(
[]previousResponseStoreSource{
{store: connectionStore, connectionLocal: true},
{store: 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"}),
)
})
})
+59 -4
View File
@@ -164,6 +164,12 @@ func NewResponseStore(ttl time.Duration) *ResponseStore {
// Store stores a response with its request and items
func (s *ResponseStore) Store(responseID string, request *schema.OpenResponsesRequest, response *schema.ORResponseResource) {
s.StoreOwned(responseID, request, response, "")
}
// StoreOwned stores a response with its owner set before the response becomes
// visible locally or through distributed replication.
func (s *ResponseStore) StoreOwned(responseID string, request *schema.OpenResponsesRequest, response *schema.ORResponseResource, owner string) {
s.mu.Lock()
// Build item index for quick lookup
@@ -182,6 +188,7 @@ func (s *ResponseStore) Store(responseID string, request *schema.OpenResponsesRe
StoredAt: time.Now(),
ExpiresAt: nil,
droppedThrough: -1,
Owner: owner,
}
// Set expiration if TTL is configured
@@ -368,6 +375,12 @@ func (s *ResponseStore) Count() int {
// StoreBackground stores a background response with cancel function and optional streaming support
func (s *ResponseStore) StoreBackground(responseID string, request *schema.OpenResponsesRequest, response *schema.ORResponseResource, cancelFunc context.CancelFunc, streamEnabled bool) {
s.StoreBackgroundOwned(responseID, request, response, cancelFunc, streamEnabled, "")
}
// StoreBackgroundOwned stores a background response with its owner set before
// the response becomes visible locally or through distributed replication.
func (s *ResponseStore) StoreBackgroundOwned(responseID string, request *schema.OpenResponsesRequest, response *schema.ORResponseResource, cancelFunc context.CancelFunc, streamEnabled bool, owner string) {
s.mu.Lock()
// Build item index for quick lookup
@@ -391,6 +404,7 @@ func (s *ResponseStore) StoreBackground(responseID string, request *schema.OpenR
IsBackground: true,
EventsChan: make(chan struct{}, 100), // Buffered channel for event notifications
droppedThrough: -1,
Owner: owner,
}
// Set expiration if TTL is configured
@@ -465,6 +479,17 @@ func (s *ResponseStore) UpdateResponse(responseID string, response *schema.ORRes
// AppendEvent appends a streaming event to the buffer for resume support
func (s *ResponseStore) AppendEvent(responseID string, event *schema.ORStreamEvent) error {
return s.appendEvent(responseID, event, false)
}
// AppendEventNext atomically assigns the sequence number immediately after the
// last buffered or evicted event, then appends the event. This is used for
// terminal events synthesized outside the normal stream producer.
func (s *ResponseStore) AppendEventNext(responseID string, event *schema.ORStreamEvent) error {
return s.appendEvent(responseID, event, true)
}
func (s *ResponseStore) appendEvent(responseID string, event *schema.ORStreamEvent, assignNext bool) error {
s.mu.RLock()
stored, exists := s.responses[responseID]
s.mu.RUnlock()
@@ -473,13 +498,23 @@ func (s *ResponseStore) AppendEvent(responseID string, event *schema.ORStreamEve
return fmt.Errorf("response not found: %s", responseID)
}
// Serialize the event
stored.mu.Lock()
defer stored.mu.Unlock()
if assignNext {
event.SequenceNumber = stored.droppedThrough + 1
if n := len(stored.StreamEvents); n > 0 {
event.SequenceNumber = stored.StreamEvents[n-1].SequenceNumber + 1
}
}
// Serialize while holding the response lock so the assigned sequence and
// append remain one atomic operation with respect to other producers.
data, err := json.Marshal(event)
if err != nil {
return fmt.Errorf("failed to marshal event: %w", err)
}
stored.mu.Lock()
stored.StreamEvents = append(stored.StreamEvents, StreamedEvent{
SequenceNumber: event.SequenceNumber,
EventType: event.Type,
@@ -504,8 +539,6 @@ func (s *ResponseStore) AppendEvent(responseID string, event *schema.ORStreamEve
stored.StreamEvents[0].Data = nil
stored.StreamEvents = stored.StreamEvents[1:]
}
stored.mu.Unlock()
// Notify any subscribers of new event
select {
case stored.EventsChan <- struct{}{}:
@@ -516,6 +549,28 @@ func (s *ResponseStore) AppendEvent(responseID string, event *schema.ORStreamEve
return nil
}
// ClearEvents releases the resume buffer while retaining the stored request
// and final response needed by previous_response_id continuation.
func (s *ResponseStore) ClearEvents(responseID string) error {
s.mu.RLock()
stored, exists := s.responses[responseID]
s.mu.RUnlock()
if !exists {
return fmt.Errorf("response not found: %s", responseID)
}
stored.mu.Lock()
for i := range stored.StreamEvents {
stored.StreamEvents[i].Data = nil
}
stored.StreamEvents = nil
stored.streamBytes = 0
stored.droppedThrough = -1
stored.mu.Unlock()
return nil
}
// GetEventsAfter returns all events with sequence number greater than startingAfter
func (s *ResponseStore) GetEventsAfter(responseID string, startingAfter int) ([]StreamedEvent, error) {
s.mu.RLock()
+286 -60
View File
@@ -3,6 +3,7 @@ package openresponses
import (
"context"
"encoding/json"
"errors"
"fmt"
"net/http"
"sync"
@@ -22,10 +23,14 @@ import (
)
const (
wsMaxMessageSize = 10 * 1024 * 1024 // 10MB
wsConnectionLimit = 60 * time.Minute
wsMaxMessageSize = 10 * 1024 * 1024 // 10MB
wsConnectionLimit = 60 * time.Minute
wsMaxConnectionLocalResponses = 128
wsMaxConnectionLocalBytes = 64 << 20 // 64 MiB
)
var errWSStreamEndedWithoutTerminal = errors.New("response stream ended without a terminal event")
var wsUpgrader = websocket.Upgrader{
CheckOrigin: func(r *http.Request) bool {
return true
@@ -38,12 +43,27 @@ type lockedConn struct {
sync.Mutex
}
type wsEventWriter interface {
writeJSON(any) error
writeTerminalJSON(any, func()) error
}
func (lc *lockedConn) writeJSON(v any) error {
lc.Lock()
defer lc.Unlock()
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://<host>:<port>/v1/responses and send response.create messages.
// Events are streamed back over the WebSocket connection instead of SSE.
@@ -54,6 +74,7 @@ func WebSocketEndpoint(application *application.Application) echo.HandlerFunc {
appConfig := application.ApplicationConfig()
return func(c echo.Context) error {
owner := ownerFromContext(c)
ws, err := wsUpgrader.Upgrade(c.Response(), c.Request(), nil)
if err != nil {
return err
@@ -75,13 +96,23 @@ func WebSocketEndpoint(application *application.Application) echo.HandlerFunc {
xlog.Debug("WebSocket Responses connection established", "address", ws.RemoteAddr().String())
handleWebSocketConnection(connCtx, conn, cl, ml, evaluator, appConfig)
handleWebSocketConnection(connCtx, conn, owner, cl, ml, evaluator, appConfig)
return nil
}
}
// 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) {
func handleWebSocketConnection(connCtx context.Context, conn *lockedConn, owner string, 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
@@ -128,10 +159,20 @@ func handleWebSocketConnection(connCtx context.Context, conn *lockedConn, cl *co
sendWSError(conn, "invalid_request", "a response is already in progress on this connection", "")
continue
}
shouldStore := wsMsg.Store == nil || *wsMsg.Store
if !shouldStore && !connectionStoreCanAccept(connectionStore, len(msgBytes), wsMaxConnectionLocalResponses, wsMaxConnectionLocalBytes) {
inflight.Unlock()
sendWSErrorEvent(conn, "connection_store_limit_reached", "connection-local response history limit reached; reconnect to start a new local history", "store")
continue
}
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, owner, wsMsg.Generate, &wsMsg.OpenResponsesRequest, cl, ml, evaluator, appConfig)
}()
}
}
@@ -140,12 +181,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(), owner string, 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 +200,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 +209,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 +220,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 +230,87 @@ func handleWSResponseCreate(connCtx context.Context, conn *lockedConn, input *sc
shouldStore = false
}
// Handle previous_response_id
store := globalStore
if !shouldStore {
store = connectionStore
}
// Resolve and authorize the complete continuation chain before prewarm or
// generation. A stored response cannot depend on connection-local history,
// because that private ancestor disappears when this socket disconnects.
var messages []schema.Message
if input.PreviousResponseID != "" {
stored, err := store.Get(input.PreviousResponseID)
previousMessages, usedConnectionLocal, err := resolvePreviousResponseMessagesFromSources(
[]previousResponseStoreSource{
{store: connectionStore, connectionLocal: true},
{store: globalStore},
},
input.PreviousResponseID,
cfg,
owner,
)
if err != nil {
sendWSErrorEvent(conn, "previous_response_not_found",
fmt.Sprintf("previous response not found: %s", input.PreviousResponseID), "previous_response_id")
if notFound, ok := err.(*previousResponseNotFoundError); ok {
failEvent("previous_response_not_found", notFound.Error(), "previous_response_id")
return
}
fail("invalid_request", err.Error(), "")
return
}
previousInputMessages, err := convertORInputToMessages(stored.Request.Input, cfg)
if err != nil {
sendWSError(conn, "invalid_request", fmt.Sprintf("failed to convert previous input: %v", err), "")
if shouldStore && usedConnectionLocal {
failEvent("invalid_store_transition", "store=true cannot persist a response that depends on connection-local store=false history; use store=false or start a new stored chain", "store")
return
}
messages = previousMessages
}
previousOutputMessages, err := convertOROutputItemsToMessages(stored.Response.Output)
if err != nil {
sendWSError(conn, "invalid_request", fmt.Sprintf("failed to convert previous response: %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.StoreBackgroundOwned(responseID, input, responseCreated, reqCancel, true, owner)
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,
})
messages = previousInputMessages
messages = append(messages, previousOutputMessages...)
processDone := make(chan struct{})
close(processDone)
lastSequence, forwardErr := forwardEvents(reqCtx, conn, store, responseID, processDone, release)
if !shouldStore {
if err := store.ClearEvents(responseID); err != nil {
xlog.Warn("WebSocket Responses: failed to clear local prewarm events", "response_id", responseID, "error", err)
}
}
if forwardErr != nil && connCtx.Err() == nil {
if err := writeWSForwardingFailure(conn, release, responseID, createdAt, lastSequence, input, shouldStore, forwardErr); err != nil {
xlog.Debug("WebSocket Responses: failed to write forwarding failure", "response_id", responseID, "error", err)
}
}
return
}
// 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...)
@@ -300,7 +398,7 @@ func handleWSResponseCreate(connCtx context.Context, conn *lockedConn, input *sc
// Use the background stream infrastructure: store the request as a background task,
// process it via handleBackgroundStream, and forward buffered events over WebSocket.
queuedResponse := buildORResponse(responseID, createdAt, nil, schema.ORStatusQueued, input, []schema.ORItemField{}, nil, shouldStore)
store.StoreBackground(responseID, input, queuedResponse, reqCancel, true)
store.StoreBackgroundOwned(responseID, input, queuedResponse, reqCancel, true, owner)
// Start processing in a goroutine
processDone := make(chan struct{})
@@ -308,22 +406,27 @@ 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()
store.UpdateStatus(responseID, schema.ORStatusFailed, &now)
// Buffer an error event so the client sees the failure
// Allocate the failure event after the producer's last event. Using a
// default sequence of zero can make clients miss the terminal event.
failedResponse := buildORResponse(responseID, createdAt, &now, schema.ORStatusFailed, input, []schema.ORItemField{}, nil, shouldStore)
bufferEvent(store, responseID, &schema.ORStreamEvent{
failureEvent := &schema.ORStreamEvent{
Type: "response.failed",
Response: failedResponse,
Error: &schema.ORErrorPayload{
Type: "server_error",
Message: bgErr.Error(),
},
})
}
normalizeORStreamEvent(failureEvent)
if err := store.AppendEventNext(responseID, failureEvent); err != nil {
xlog.Error("WebSocket Responses: failed to buffer failure event", "response_id", responseID, "error", err)
}
return
}
if finalResponse != nil {
@@ -332,72 +435,166 @@ 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)
lastSequence, forwardErr := forwardEvents(reqCtx, conn, store, responseID, processDone, release)
if forwardErr != nil {
reqCancel()
select {
case <-processDone:
case <-connCtx.Done():
}
}
if !shouldStore {
if err := store.ClearEvents(responseID); err != nil {
xlog.Warn("WebSocket Responses: failed to clear local response events", "response_id", responseID, "error", err)
}
}
if forwardErr != nil && connCtx.Err() == nil {
xlog.Error("WebSocket Responses: event forwarding failed", "response_id", responseID, "error", forwardErr)
if err := writeWSForwardingFailure(conn, release, responseID, createdAt, lastSequence, input, shouldStore, forwardErr); err != nil {
xlog.Debug("WebSocket Responses: failed to write forwarding failure", "response_id", responseID, "error", err)
}
}
}
// 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) {
// The returned sequence is the last event successfully delivered to the client.
func forwardEvents(ctx context.Context, conn wsEventWriter, store *ResponseStore, responseID string, done <-chan struct{}, release func()) (int, error) {
eventsChan, err := store.GetEventsChan(responseID)
if err != nil {
return
return -1, err
}
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
return lastSeq, err
}
for _, event := range events {
var parsed schema.ORStreamEvent
if err := json.Unmarshal(event.Data, &parsed); err != nil {
continue
return lastSeq, fmt.Errorf("failed to decode buffered response event: %w", err)
}
if err := conn.writeJSON(&parsed); err != nil {
return
terminal, err := writeEvent(&parsed)
if err != nil {
return lastSeq, err
}
if terminal {
return event.SequenceNumber, nil
}
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 {
var parsed schema.ORStreamEvent
if err := json.Unmarshal(event.Data, &parsed); err != nil {
continue
}
if err := conn.writeJSON(&parsed); err != nil {
return
}
if err != nil {
return lastSeq, err
}
for _, event := range finalEvents {
var parsed schema.ORStreamEvent
if err := json.Unmarshal(event.Data, &parsed); err != nil {
return lastSeq, fmt.Errorf("failed to decode buffered response event: %w", err)
}
terminal, err := writeEvent(&parsed)
if err != nil {
return lastSeq, err
}
if terminal {
return event.SequenceNumber, nil
}
lastSeq = event.SequenceNumber
}
// Clean up non-stored responses from the cache
if !shouldStore {
store.Delete(responseID)
}
return
return lastSeq, errWSStreamEndedWithoutTerminal
default:
}
// Wait for new events, completion, or context cancellation
// Wait for new events, completion, or context cancellation.
select {
case <-ctx.Done():
return
return lastSeq, ctx.Err()
case <-done:
// Will drain in next iteration
// Will drain in next iteration.
case <-eventsChan:
// New events available
// New events available.
}
}
}
func connectionStoreUsage(store *ResponseStore) (count, bytes int) {
store.mu.RLock()
responses := make([]*StoredResponse, 0, len(store.responses))
for _, stored := range store.responses {
responses = append(responses, stored)
}
store.mu.RUnlock()
for _, stored := range responses {
stored.mu.RLock()
requestData, _ := json.Marshal(stored.Request)
responseData, _ := json.Marshal(stored.Response)
streamBytes := stored.streamBytes
if stored.Response != nil && isORTerminalStatus(stored.Response.Status) {
// The terminal event has already released this response's in-flight
// guard. Exclude its soon-to-be-cleared delivery buffer so the next
// request cannot race cleanup and receive a false quota rejection.
streamBytes = 0
}
bytes += len(requestData) + len(responseData) + streamBytes
stored.mu.RUnlock()
}
return len(responses), bytes
}
func isORTerminalStatus(status string) bool {
return status == schema.ORStatusCompleted || status == schema.ORStatusFailed ||
status == schema.ORStatusIncomplete || status == schema.ORStatusCancelled
}
func connectionStoreCanAccept(store *ResponseStore, incomingBytes, maxResponses, maxBytes int) bool {
count, bytes := connectionStoreUsage(store)
if maxResponses > 0 && count >= maxResponses {
return false
}
return maxBytes <= 0 || bytes+incomingBytes <= maxBytes
}
func writeWSForwardingFailure(conn wsEventWriter, release func(), responseID string, createdAt int64, lastSequence int, input *schema.OpenResponsesRequest, shouldStore bool, forwardingErr error) error {
now := time.Now().Unix()
failedResponse := buildORResponse(responseID, createdAt, &now, schema.ORStatusFailed, input, []schema.ORItemField{}, nil, shouldStore)
event := &schema.ORStreamEvent{
Type: "response.failed",
SequenceNumber: lastSequence + 1,
Response: failedResponse,
Error: &schema.ORErrorPayload{
Type: "server_error",
Message: fmt.Sprintf("response stream failed: %v", forwardingErr),
},
}
normalizeORStreamEvent(event)
return conn.writeTerminalJSON(event, release)
}
func sendWSError(conn *lockedConn, errType, message, param string) {
event := schema.ORStreamEvent{
Type: "error",
@@ -410,6 +607,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 +633,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)
}
}
@@ -0,0 +1,252 @@
package openresponses
import (
"context"
"encoding/json"
"errors"
"time"
"github.com/mudler/LocalAI/core/config"
"github.com/mudler/LocalAI/core/schema"
. "github.com/onsi/ginkgo/v2"
. "github.com/onsi/gomega"
)
type recordingWSEventWriter struct {
events []*schema.ORStreamEvent
onWrite func()
}
func (w *recordingWSEventWriter) writeJSON(v any) error {
event, ok := v.(*schema.ORStreamEvent)
Expect(ok).To(BeTrue())
w.events = append(w.events, event)
if w.onWrite != nil {
w.onWrite()
}
return nil
}
func (w *recordingWSEventWriter) writeTerminalJSON(v any, release func()) error {
release()
return w.writeJSON(v)
}
var _ = Describe("WebSocket Responses", func() {
It("allocates a failure event after the last buffered sequence", func() {
store := NewResponseStore(0)
store.StoreBackground(
"resp_failure",
&schema.OpenResponsesRequest{Model: "test-model", Input: "hello"},
&schema.ORResponseResource{ID: "resp_failure", Status: schema.ORStatusInProgress},
func() {},
true,
)
Expect(store.AppendEvent("resp_failure", &schema.ORStreamEvent{
Type: "response.in_progress",
SequenceNumber: 7,
})).To(Succeed())
failure := &schema.ORStreamEvent{Type: "response.failed"}
Expect(store.AppendEventNext("resp_failure", failure)).To(Succeed())
Expect(failure.SequenceNumber).To(Equal(8))
events, err := store.GetEventsAfter("resp_failure", 7)
Expect(err).NotTo(HaveOccurred())
Expect(events).To(HaveLen(1))
Expect(events[0].SequenceNumber).To(Equal(8))
var persisted schema.ORStreamEvent
Expect(json.Unmarshal(events[0].Data, &persisted)).To(Succeed())
Expect(persisted.SequenceNumber).To(Equal(8))
Expect(persisted.Type).To(Equal("response.failed"))
})
It("propagates a lost event offset instead of silently ending forwarding", func() {
store := NewResponseStore(0)
store.maxStreamEvents = 1
store.StoreBackground(
"resp_gap",
&schema.OpenResponsesRequest{Model: "test-model", Input: "hello"},
&schema.ORResponseResource{ID: "resp_gap", Status: schema.ORStatusInProgress},
func() {},
true,
)
Expect(store.AppendEvent("resp_gap", &schema.ORStreamEvent{Type: "response.created", SequenceNumber: 0})).To(Succeed())
Expect(store.AppendEvent("resp_gap", &schema.ORStreamEvent{Type: "response.in_progress", SequenceNumber: 1})).To(Succeed())
done := make(chan struct{})
lastSequence, err := forwardEvents(context.Background(), &recordingWSEventWriter{}, store, "resp_gap", done, func() {})
Expect(errors.Is(err, ErrOffsetLost)).To(BeTrue())
Expect(lastSequence).To(Equal(-1))
})
It("reports the last delivered sequence when forwarding fails after partial delivery", func() {
store := NewResponseStore(0)
store.maxStreamEvents = 1
store.StoreBackground(
"resp_partial_gap",
&schema.OpenResponsesRequest{Model: "test-model", Input: "hello"},
&schema.ORResponseResource{ID: "resp_partial_gap", Status: schema.ORStatusInProgress},
func() {},
true,
)
Expect(store.AppendEvent("resp_partial_gap", &schema.ORStreamEvent{Type: "response.created", SequenceNumber: 0})).To(Succeed())
appended := false
writer := &recordingWSEventWriter{}
writer.onWrite = func() {
if appended {
return
}
appended = true
Expect(store.AppendEvent("resp_partial_gap", &schema.ORStreamEvent{Type: "response.in_progress", SequenceNumber: 1})).To(Succeed())
Expect(store.AppendEvent("resp_partial_gap", &schema.ORStreamEvent{Type: "response.output_item.added", SequenceNumber: 2})).To(Succeed())
}
lastSequence, err := forwardEvents(context.Background(), writer, store, "resp_partial_gap", make(chan struct{}), func() {})
Expect(errors.Is(err, ErrOffsetLost)).To(BeTrue())
Expect(lastSequence).To(Equal(0))
Expect(writer.events).To(HaveLen(1))
})
It("terminates a forwarding failure with the next delivered sequence", func() {
writer := &recordingWSEventWriter{}
released := false
err := writeWSForwardingFailure(
writer,
func() { released = true },
"resp_forwarding_failure",
time.Now().Unix(),
7,
&schema.OpenResponsesRequest{Model: "test-model"},
false,
errors.New("offset lost"),
)
Expect(err).NotTo(HaveOccurred())
Expect(released).To(BeTrue())
Expect(writer.events).To(HaveLen(1))
Expect(writer.events[0].Type).To(Equal("response.failed"))
Expect(writer.events[0].SequenceNumber).To(Equal(8))
Expect(writer.events[0].Response.Status).To(Equal(schema.ORStatusFailed))
})
It("does not resolve a globally stored response owned by another caller", func() {
globalStore := NewResponseStore(0)
globalStore.StoreOwned("resp_private", &schema.OpenResponsesRequest{Input: "secret"}, &schema.ORResponseResource{
ID: "resp_private",
Output: []schema.ORItemField{},
}, "user-a")
_, _, err := resolvePreviousResponseMessagesFromSources(
[]previousResponseStoreSource{{store: globalStore}},
"resp_private",
&config.ModelConfig{},
"user-b",
)
var notFound *previousResponseNotFoundError
Expect(errors.As(err, &notFound)).To(BeTrue())
})
It("checks ownership on every hop of a stored continuation chain", func() {
globalStore := NewResponseStore(0)
globalStore.Store("resp_ancestor", &schema.OpenResponsesRequest{Input: "secret"}, &schema.ORResponseResource{
ID: "resp_ancestor",
Output: []schema.ORItemField{},
})
globalStore.SetOwner("resp_ancestor", "user-a")
globalStore.Store("resp_head", &schema.OpenResponsesRequest{PreviousResponseID: "resp_ancestor", Input: "mine"}, &schema.ORResponseResource{
ID: "resp_head",
Output: []schema.ORItemField{},
})
globalStore.SetOwner("resp_head", "user-b")
_, _, err := resolvePreviousResponseMessagesFromSources(
[]previousResponseStoreSource{{store: globalStore}},
"resp_head",
&config.ModelConfig{},
"user-b",
)
var notFound *previousResponseNotFoundError
Expect(errors.As(err, &notFound)).To(BeTrue())
Expect(notFound.ResponseID).To(Equal("resp_ancestor"))
})
It("reports when continuation history uses the connection-local store", func() {
connectionStore := NewResponseStore(0)
connectionStore.Store("resp_local", &schema.OpenResponsesRequest{Input: "private"}, &schema.ORResponseResource{
ID: "resp_local",
Output: []schema.ORItemField{},
})
messages, usedConnectionLocal, err := resolvePreviousResponseMessagesFromSources(
[]previousResponseStoreSource{{store: connectionStore, connectionLocal: true}},
"resp_local",
&config.ModelConfig{},
"",
)
Expect(err).NotTo(HaveOccurred())
Expect(messages).To(HaveLen(1))
Expect(usedConnectionLocal).To(BeTrue())
})
It("reports a connection-local ancestor behind a globally stored head", func() {
connectionStore := NewResponseStore(0)
globalStore := NewResponseStore(0)
connectionStore.Store("resp_local", &schema.OpenResponsesRequest{Input: "private"}, &schema.ORResponseResource{
ID: "resp_local",
Output: []schema.ORItemField{},
})
globalStore.Store("resp_global", &schema.OpenResponsesRequest{PreviousResponseID: "resp_local", Input: "child"}, &schema.ORResponseResource{
ID: "resp_global",
Output: []schema.ORItemField{},
})
_, usedConnectionLocal, err := resolvePreviousResponseMessagesFromSources(
[]previousResponseStoreSource{
{store: connectionStore, connectionLocal: true},
{store: globalStore},
},
"resp_global",
&config.ModelConfig{},
"",
)
Expect(err).NotTo(HaveOccurred())
Expect(usedConnectionLocal).To(BeTrue())
})
It("clears forwarded events without deleting continuation state", func() {
store := NewResponseStore(0)
store.StoreBackground(
"resp_local",
&schema.OpenResponsesRequest{Model: "test-model", Input: "hello"},
&schema.ORResponseResource{ID: "resp_local", Status: schema.ORStatusCompleted},
func() {},
true,
)
Expect(store.AppendEvent("resp_local", &schema.ORStreamEvent{Type: "response.completed", SequenceNumber: 0})).To(Succeed())
_, usageBeforeCleanup := connectionStoreUsage(store)
Expect(store.ClearEvents("resp_local")).To(Succeed())
_, usageAfterCleanup := connectionStoreUsage(store)
Expect(usageBeforeCleanup).To(Equal(usageAfterCleanup), "terminal stream buffers should not affect admission while cleanup finishes")
events, err := store.GetEventsAfter("resp_local", -1)
Expect(err).NotTo(HaveOccurred())
Expect(events).To(BeEmpty())
stored, err := store.Get("resp_local")
Expect(err).NotTo(HaveOccurred())
Expect(stored.Request.Input).To(Equal("hello"))
Expect(stored.Response.Status).To(Equal(schema.ORStatusCompleted))
})
It("rejects new local history when the connection budget is exhausted", func() {
store := NewResponseStore(0)
store.Store("resp_local", &schema.OpenResponsesRequest{Input: "1234567890"}, &schema.ORResponseResource{ID: "resp_local"})
count, size := connectionStoreUsage(store)
Expect(count).To(Equal(1))
Expect(size).To(BeNumerically(">", 0))
Expect(connectionStoreCanAccept(store, 1, 2, size+1)).To(BeTrue())
Expect(connectionStoreCanAccept(store, 1, 1, size+1)).To(BeFalse())
Expect(connectionStoreCanAccept(store, 2, 2, size+1)).To(BeFalse())
})
})
+9 -5
View File
@@ -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"`
+55
View File
@@ -240,6 +240,61 @@ curl http://localhost:8080/v1/responses \
}'
```
#### WebSocket Responses
Connect to `ws://localhost:8080/v1/responses` (or `wss://` when TLS is
enabled) and send `response.create` messages over the WebSocket. Only one
response may be in progress on a connection. Wait for `response.completed` or
`response.failed` for the active response before sending the next
`response.create`. An `error` that rejects an invalid or additional
`response.create` applies only to that rejected message; it does not terminate
a response that is already in progress.
Set the WebSocket-only `generate` field to `false` to prepare a request without
running inference:
```json
{
"type": "response.create",
"model": "ggml-koala-7b-model-q4_0-r2.bin",
"generate": false,
"store": false,
"input": "Say this is a test!"
}
```
LocalAI emits `response.created` followed by `response.completed` with no
generated output. The completed response still has an ID that can continue the
prepared request:
```json
{
"type": "response.create",
"model": "ggml-koala-7b-model-q4_0-r2.bin",
"store": false,
"previous_response_id": "resp_abc123",
"input": []
}
```
Response IDs created with `store: false` are available only on the WebSocket
connection that created them and are removed when that connection closes. A
`previous_response_id` chain can contain multiple responses; LocalAI replays the
complete conversation from the oldest response through the referenced response
before appending the new input. Keep `store: false` for every descendant of a
connection-local response. LocalAI rejects a `store: true` response that depends
on connection-local history because the resulting stored chain would be broken
after the WebSocket closes.
After delivery, LocalAI discards buffered stream events for `store: false`
responses but retains their request and final response for continuation. A
connection admits at most 128 local responses and accepts a new local request
only while its current serialized history plus that request is at most 64 MiB.
When either admission threshold is reached, LocalAI returns
`connection_store_limit_reached`; reconnect to start a new local history. With
authentication enabled, globally stored response IDs can be continued only by
the identity that created them.
#### Request Parameters
| Parameter | Type | Required | Description |
+152 -1
View File
@@ -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()
@@ -199,6 +267,89 @@ var _ = Describe("WebSocket Responses API E2E Tests", Label("WebSocket"), func()
})
Context("Error handling", func() {
It("rejects a concurrent create without terminating the active response", func() {
conn, err := dialWS()
Expect(err).ToNot(HaveOccurred())
defer func() { _ = conn.Close() }()
Expect(conn.WriteJSON(map[string]any{
"type": "response.create",
"model": "mock-model",
"input": "MOCK_SLOW_STREAM",
})).To(Succeed())
Expect(conn.WriteJSON(map[string]any{
"type": "response.create",
"model": "mock-model",
"input": "must be rejected",
})).To(Succeed())
sawRejection := false
sawCompletion := false
for !sawRejection || !sawCompletion {
ev, err := readEvent(conn)
Expect(err).NotTo(HaveOccurred())
switch ev.Type {
case "error":
Expect(ev.Error).NotTo(BeNil())
Expect(ev.Error.Message).To(ContainSubstring("already in progress"))
sawRejection = true
case "response.completed":
sawCompletion = true
}
}
})
It("emits a sequenced terminal failure when inference fails mid-stream", func() {
conn, err := dialWS()
Expect(err).ToNot(HaveOccurred())
defer func() { _ = conn.Close() }()
Expect(conn.WriteJSON(map[string]any{
"type": "response.create",
"model": "mock-model",
"store": false,
"input": "MOCK_ERROR_MIDSTREAM",
})).To(Succeed())
events := readAllEvents(conn)
Expect(events).ToNot(BeEmpty())
Expect(events[len(events)-1].Type).To(Equal("response.failed"))
for i := 1; i < len(events); i++ {
Expect(events[i].SequenceNumber).To(BeNumerically(">", events[i-1].SequenceNumber))
}
})
It("rejects storing a response whose history is connection-local", func() {
conn, err := dialWS()
Expect(err).ToNot(HaveOccurred())
defer func() { _ = conn.Close() }()
Expect(conn.WriteJSON(map[string]any{
"type": "response.create",
"model": "mock-model",
"store": false,
"generate": false,
"input": "private warmup",
})).To(Succeed())
warmupEvents := readAllEvents(conn)
Expect(warmupEvents).To(HaveLen(2))
var warmupResp wsResponseBody
Expect(json.Unmarshal(warmupEvents[1].Response, &warmupResp)).To(Succeed())
Expect(conn.WriteJSON(map[string]any{
"type": "response.create",
"model": "mock-model",
"store": true,
"previous_response_id": warmupResp.ID,
"input": "make this global",
})).To(Succeed())
ev, err := readEvent(conn)
Expect(err).NotTo(HaveOccurred())
Expect(ev.Type).To(Equal("error"))
Expect(ev.Error).NotTo(BeNil())
Expect(ev.Error.Code).To(Equal("invalid_store_transition"))
})
It("returns error for previous_response_not_found", func() {
conn, err := dialWS()
Expect(err).ToNot(HaveOccurred())
+8
View File
@@ -14,6 +14,7 @@ import (
"strconv"
"strings"
"sync"
"time"
"github.com/mudler/LocalAI/pkg/grpc/grpcerrors"
pb "github.com/mudler/LocalAI/pkg/grpc/proto"
@@ -270,6 +271,13 @@ func (m *MockBackend) PredictStream(in *pb.PredictOptions, stream pb.Backend_Pre
}
return fmt.Errorf("mock backend stream error: simulated mid-stream failure")
}
if strings.Contains(in.Prompt, "MOCK_SLOW_STREAM") {
select {
case <-time.After(500 * time.Millisecond):
case <-stream.Context().Done():
return stream.Context().Err()
}
}
// Simulate C++ autoparser behavior: tool calls delivered via ChatDeltas
// with empty message (autoparser clears raw message during parsing).