Skip to content

Commit b1b2fbb

Browse files
committed
fix(realtime): Improve tool call handling and error reporting
- Refactor Model interface to accept []types.ToolUnion and *types.ToolChoiceUnion instead of JSON strings, eliminating unnecessary marshal/unmarshal cycles - Fix Parameters field handling: support both map[string]any and JSON string formats - Add PredictConfig() method to Model interface for accessing model configuration - Add comprehensive debug logging for tool call parsing and function config - Add missing return statement after prediction error (critical bug fix) - Add warning logs for NoAction function argument parsing failures - Improve error visibility throughout generateResponse function 💘 Generated with Crush Assisted-by: Claude Sonnet 4.5 via Crush <crush@charm.land> Signed-off-by: Richard Palethorpe <io@richiejp.com>
1 parent a79a15e commit b1b2fbb

2 files changed

Lines changed: 123 additions & 24 deletions

File tree

‎core/http/endpoints/openai/realtime.go‎

Lines changed: 25 additions & 20 deletions
Original file line numberDiff line numberDiff line change
@@ -41,10 +41,10 @@ const (
4141

4242
// Session represents a single WebSocket connection and its state
4343
type Session struct {
44-
ID string
45-
TranscriptionOnly bool
44+
ID string
45+
TranscriptionOnly bool
4646
// The pipeline or any-to-any model name (full realtime mode)
47-
Model string
47+
Model string
4848
// The voice may be a TTS model name or a parameter passed to a TTS model
4949
Voice string
5050
TurnDetection *types.TurnDetectionUnion // "server_vad", "semantic_vad" or "none"
@@ -58,7 +58,7 @@ type Session struct {
5858
DefaultConversationID string
5959
ModelInterface Model
6060
// The pipeline model config or the config for an any-to-any model
61-
ModelConfig *config.ModelConfig
61+
ModelConfig *config.ModelConfig
6262
}
6363

6464
func (s *Session) FromClient(session *types.SessionUnion) {
@@ -121,8 +121,9 @@ var sessionLock sync.Mutex
121121
type Model interface {
122122
VAD(ctx context.Context, request *schema.VADRequest) (*schema.VADResponse, error)
123123
Transcribe(ctx context.Context, audio, language string, translate bool, diarize bool, prompt string) (*schema.TranscriptionResult, error)
124-
Predict(ctx context.Context, messages schema.Messages, images, videos, audios []string, tokenCallback func(string, backend.TokenUsage) bool, tools string, toolChoice string, logprobs *int, topLogprobs *int, logitBias map[string]float64) (func() (backend.LLMResponse, error), error)
124+
Predict(ctx context.Context, messages schema.Messages, images, videos, audios []string, tokenCallback func(string, backend.TokenUsage) bool, tools []types.ToolUnion, toolChoice *types.ToolChoiceUnion, logprobs *int, topLogprobs *int, logitBias map[string]float64) (func() (backend.LLMResponse, error), error)
125125
TTS(ctx context.Context, text, voice, language string) (string, *proto.Result, error)
126+
PredictConfig() *config.ModelConfig
126127
}
127128

128129
var upgrader = websocket.Upgrader{
@@ -765,7 +766,7 @@ func commitUtterance(ctx context.Context, utt []byte, session *Session, conv *Co
765766
}
766767

767768
if !session.TranscriptionOnly {
768-
generateResponse(session.ModelConfig, session, utt, transcript, conv, c, websocket.TextMessage)
769+
generateResponse(session, utt, transcript, conv, c, websocket.TextMessage)
769770
}
770771
}
771772

@@ -790,9 +791,11 @@ func runVAD(ctx context.Context, session *Session, adata []int16) ([]schema.VADS
790791
}
791792

792793
// Function to generate a response based on the conversation
793-
func generateResponse(config *config.ModelConfig, session *Session, utt []byte, transcript string, conv *Conversation, c *websocket.Conn, mt int) {
794+
func generateResponse(session *Session, utt []byte, transcript string, conv *Conversation, c *websocket.Conn, mt int) {
794795
xlog.Debug("Generating realtime response...")
795796

797+
config := session.ModelInterface.PredictConfig()
798+
796799
item := types.MessageItemUnion{
797800
User: &types.MessageItemUser{
798801
ID: generateItemID(),
@@ -881,19 +884,7 @@ func generateResponse(config *config.ModelConfig, session *Session, utt []byte,
881884
},
882885
})
883886

884-
toolsJSON := ""
885-
if len(session.Tools) > 0 {
886-
b, _ := json.Marshal(session.Tools)
887-
toolsJSON = string(b)
888-
}
889-
890-
toolChoiceJSON := ""
891-
if session.ToolChoice != nil {
892-
b, _ := json.Marshal(session.ToolChoice)
893-
toolChoiceJSON = string(b)
894-
}
895-
896-
predFunc, err := session.ModelInterface.Predict(context.TODO(), conversationHistory, nil, nil, nil, nil, toolsJSON, toolChoiceJSON, nil, nil, nil)
887+
predFunc, err := session.ModelInterface.Predict(context.TODO(), conversationHistory, nil, nil, nil, nil, session.Tools, session.ToolChoice, nil, nil, nil)
897888
if err != nil {
898889
sendError(c, "inference_failed", fmt.Sprintf("backend error: %v", err), "", item.Assistant.ID)
899890
return
@@ -902,8 +893,11 @@ func generateResponse(config *config.ModelConfig, session *Session, utt []byte,
902893
pred, err := predFunc()
903894
if err != nil {
904895
sendError(c, "prediction_failed", fmt.Sprintf("backend error: %v", err), "", item.Assistant.ID)
896+
return
905897
}
906898

899+
xlog.Debug("Function config for parsing", "function_name_key", config.FunctionsConfig.FunctionNameKey, "function_arguments_key", config.FunctionsConfig.FunctionArgumentsKey)
900+
907901
rawResponse := pred.Response
908902
if config.TemplateConfig.ReplyPrefix != "" {
909903
rawResponse = config.TemplateConfig.ReplyPrefix + rawResponse
@@ -916,6 +910,8 @@ func generateResponse(config *config.ModelConfig, session *Session, utt []byte,
916910
cleanedResponse := functions.CleanupLLMResult(responseWithoutReasoning, config.FunctionsConfig)
917911
toolCalls := functions.ParseFunctionCall(cleanedResponse, config.FunctionsConfig)
918912

913+
xlog.Debug("Function call parsing", "textContent", textContent, "cleanedResponse", cleanedResponse, "toolCallsCount", len(toolCalls))
914+
919915
noActionName := "answer"
920916
if config.FunctionsConfig.NoActionFunctionName != "" {
921917
noActionName = config.FunctionsConfig.NoActionFunctionName
@@ -932,15 +928,23 @@ func generateResponse(config *config.ModelConfig, session *Session, utt []byte,
932928
if m, exists := arguments["message"]; exists {
933929
if message, ok := m.(string); ok {
934930
finalSpeech = message
931+
} else {
932+
xlog.Warn("NoAction function message field is not a string", "type", fmt.Sprintf("%T", m))
935933
}
934+
} else {
935+
xlog.Warn("NoAction function missing 'message' field in arguments")
936936
}
937+
} else {
938+
xlog.Warn("Failed to unmarshal NoAction function arguments", "error", err, "arguments", arg)
937939
}
938940
if finalSpeech == "" {
939941
// Fallback if parsing failed
942+
xlog.Warn("NoAction function did not produce speech, using cleaned response as fallback")
940943
finalSpeech = cleanedResponse
941944
}
942945
} else {
943946
finalToolCalls = toolCalls
947+
xlog.Debug("Setting finalToolCalls", "count", len(finalToolCalls))
944948
if len(toolCalls) > 0 {
945949
finalSpeech = textContent
946950
} else {
@@ -1060,6 +1064,7 @@ func generateResponse(config *config.ModelConfig, session *Session, utt []byte,
10601064
}
10611065

10621066
// Handle Tool Calls
1067+
xlog.Debug("About to handle tool calls", "finalToolCallsCount", len(finalToolCalls))
10631068
for i, tc := range finalToolCalls {
10641069
toolCallID := generateItemID()
10651070
callID := "call_" + generateUniqueID() // OpenAI uses call_xyz

‎core/http/endpoints/openai/realtime_model.go‎

Lines changed: 98 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -2,10 +2,12 @@ package openai
22

33
import (
44
"context"
5+
"encoding/json"
56
"fmt"
67

78
"github.com/mudler/LocalAI/core/backend"
89
"github.com/mudler/LocalAI/core/config"
10+
"github.com/mudler/LocalAI/core/http/endpoints/openai/types"
911
"github.com/mudler/LocalAI/core/schema"
1012
"github.com/mudler/LocalAI/core/templates"
1113
"github.com/mudler/LocalAI/pkg/functions"
@@ -63,14 +65,18 @@ func (m *transcriptOnlyModel) Transcribe(ctx context.Context, audio, language st
6365
return backend.ModelTranscription(audio, language, translate, diarize, prompt, m.modelLoader, *m.TranscriptionConfig, m.appConfig)
6466
}
6567

66-
func (m *transcriptOnlyModel) Predict(ctx context.Context, messages schema.Messages, images, videos, audios []string, tokenCallback func(string, backend.TokenUsage) bool, tools string, toolChoice string, logprobs *int, topLogprobs *int, logitBias map[string]float64) (func() (backend.LLMResponse, error), error) {
68+
func (m *transcriptOnlyModel) Predict(ctx context.Context, messages schema.Messages, images, videos, audios []string, tokenCallback func(string, backend.TokenUsage) bool, tools []types.ToolUnion, toolChoice *types.ToolChoiceUnion, logprobs *int, topLogprobs *int, logitBias map[string]float64) (func() (backend.LLMResponse, error), error) {
6769
return nil, fmt.Errorf("predict operation not supported in transcript-only mode")
6870
}
6971

7072
func (m *transcriptOnlyModel) TTS(ctx context.Context, text, voice, language string) (string, *proto.Result, error) {
7173
return "", nil, fmt.Errorf("TTS not supported in transcript-only mode")
7274
}
7375

76+
func (m *transcriptOnlyModel) PredictConfig() *config.ModelConfig {
77+
return nil
78+
}
79+
7480
func (m *wrappedModel) VAD(ctx context.Context, request *schema.VADRequest) (*schema.VADResponse, error) {
7581
return backend.VAD(request, ctx, m.modelLoader, m.appConfig, *m.VADConfig)
7682
}
@@ -79,28 +85,116 @@ func (m *wrappedModel) Transcribe(ctx context.Context, audio, language string, t
7985
return backend.ModelTranscription(audio, language, translate, diarize, prompt, m.modelLoader, *m.TranscriptionConfig, m.appConfig)
8086
}
8187

82-
func (m *wrappedModel) Predict(ctx context.Context, messages schema.Messages, images, videos, audios []string, tokenCallback func(string, backend.TokenUsage) bool, tools string, toolChoice string, logprobs *int, topLogprobs *int, logitBias map[string]float64) (func() (backend.LLMResponse, error), error) {
88+
func (m *wrappedModel) Predict(ctx context.Context, messages schema.Messages, images, videos, audios []string, tokenCallback func(string, backend.TokenUsage) bool, tools []types.ToolUnion, toolChoice *types.ToolChoiceUnion, logprobs *int, topLogprobs *int, logitBias map[string]float64) (func() (backend.LLMResponse, error), error) {
8389
input := schema.OpenAIRequest{
8490
Messages: messages,
8591
}
8692

8793
var predInput string
94+
var funcs []functions.Function
8895
if !m.LLMConfig.TemplateConfig.UseTokenizerTemplate {
89-
predInput = m.evaluator.TemplateMessages(input, input.Messages, m.LLMConfig, []functions.Function{}, false)
96+
if len(tools) > 0 {
97+
for _, t := range tools {
98+
if t.Function != nil {
99+
var params map[string]any
100+
101+
switch p := t.Function.Parameters.(type) {
102+
case map[string]any:
103+
params = p
104+
case string:
105+
if err := json.Unmarshal([]byte(p), &params); err != nil {
106+
xlog.Warn("Failed to parse parameters JSON string", "error", err, "function", t.Function.Name)
107+
}
108+
}
109+
110+
funcs = append(funcs, functions.Function{
111+
Name: t.Function.Name,
112+
Description: t.Function.Description,
113+
Parameters: params,
114+
})
115+
}
116+
}
117+
}
118+
119+
predInput = m.evaluator.TemplateMessages(input, input.Messages, m.LLMConfig, funcs, len(funcs) > 0)
90120

91121
xlog.Debug("Prompt (after templating)", "prompt", predInput)
92122
if m.LLMConfig.Grammar != "" {
93123
xlog.Debug("Grammar", "grammar", m.LLMConfig.Grammar)
94124
}
95125
}
96126

97-
return backend.ModelInference(ctx, predInput, messages, images, videos, audios, m.modelLoader, m.LLMConfig, m.confLoader, m.appConfig, tokenCallback, tools, toolChoice, logprobs, topLogprobs, logitBias, )
127+
// Generate grammar for function calling if tools are provided and grammar generation is enabled
128+
shouldUseFn := len(tools) > 0 && m.LLMConfig.ShouldUseFunctions()
129+
130+
if !m.LLMConfig.FunctionsConfig.GrammarConfig.NoGrammar && shouldUseFn {
131+
// Allow the user to set custom actions via config file
132+
noActionName := "answer"
133+
noActionDescription := "use this action to answer without performing any action"
134+
135+
if m.LLMConfig.FunctionsConfig.NoActionFunctionName != "" {
136+
noActionName = m.LLMConfig.FunctionsConfig.NoActionFunctionName
137+
}
138+
if m.LLMConfig.FunctionsConfig.NoActionDescriptionName != "" {
139+
noActionDescription = m.LLMConfig.FunctionsConfig.NoActionDescriptionName
140+
}
141+
142+
noActionGrammar := functions.Function{
143+
Name: noActionName,
144+
Description: noActionDescription,
145+
Parameters: map[string]interface{}{
146+
"properties": map[string]interface{}{
147+
"message": map[string]interface{}{
148+
"type": "string",
149+
"description": "The message to reply the user with",
150+
},
151+
},
152+
},
153+
}
154+
155+
if !m.LLMConfig.FunctionsConfig.DisableNoAction {
156+
funcs = append(funcs, noActionGrammar)
157+
}
158+
159+
// Force picking one of the functions by the request
160+
if m.LLMConfig.FunctionToCall() != "" {
161+
funcs = functions.Functions(funcs).Select(m.LLMConfig.FunctionToCall())
162+
}
163+
164+
// Generate grammar from function definitions
165+
jsStruct := functions.Functions(funcs).ToJSONStructure(m.LLMConfig.FunctionsConfig.FunctionNameKey, m.LLMConfig.FunctionsConfig.FunctionNameKey)
166+
g, err := jsStruct.Grammar(m.LLMConfig.FunctionsConfig.GrammarOptions()...)
167+
if err == nil {
168+
m.LLMConfig.Grammar = g
169+
xlog.Debug("Generated grammar for function calling", "grammar", g)
170+
} else {
171+
xlog.Error("Failed generating grammar", "error", err)
172+
}
173+
}
174+
175+
var toolsJSON string
176+
if len(tools) > 0 {
177+
b, _ := json.Marshal(tools)
178+
toolsJSON = string(b)
179+
}
180+
181+
var toolChoiceJSON string
182+
if toolChoice != nil {
183+
b, _ := json.Marshal(toolChoice)
184+
toolChoiceJSON = string(b)
185+
}
186+
187+
return backend.ModelInference(ctx, predInput, messages, images, videos, audios, m.modelLoader, m.LLMConfig, m.confLoader, m.appConfig, tokenCallback, toolsJSON, toolChoiceJSON, logprobs, topLogprobs, logitBias, )
98188
}
99189

100190
func (m *wrappedModel) TTS(ctx context.Context, text, voice, language string) (string, *proto.Result, error) {
101191
return backend.ModelTTS(text, voice, language, m.modelLoader, m.appConfig, *m.TTSConfig)
102192
}
103193

194+
func (m *wrappedModel) PredictConfig() *config.ModelConfig {
195+
return m.LLMConfig
196+
}
197+
104198
func newTranscriptionOnlyModel(pipeline *config.Pipeline, cl *config.ModelConfigLoader, ml *model.ModelLoader, appConfig *config.ApplicationConfig) (Model, *config.ModelConfig, error) {
105199
cfgVAD, err := cl.LoadModelConfigFileByName(pipeline.VAD, ml.ModelPath)
106200
if err != nil {

0 commit comments

Comments
 (0)