From 6bfad9c1a6f9ecfb8388c356950cf1bbd768ea6b Mon Sep 17 00:00:00 2001 From: azhou Date: Thu, 7 May 2026 00:30:49 +0800 Subject: [PATCH 1/2] guardrail for openai non stream --- internal/guardrails/adapter/openai.go | 318 +++++++++++++++ internal/guardrails/mutate/openai.go | 371 ++++++++++++++++++ .../pipeline/openai_nonstream_response.go | 66 ++++ .../guardrails/pipeline/openai_request.go | 93 +++++ internal/server/guardrails_runtime.go | 88 +++++ internal/server/mcp_anthropic_v1_helper.go | 5 + internal/server/openai.go | 5 + internal/server/openai_chat.go | 5 + internal/server/openai_responses.go | 4 + internal/server/protocol_dispatch.go | 10 + internal/server/server.go | 3 + 11 files changed, 968 insertions(+) create mode 100644 internal/guardrails/adapter/openai.go create mode 100644 internal/guardrails/mutate/openai.go create mode 100644 internal/guardrails/pipeline/openai_nonstream_response.go create mode 100644 internal/guardrails/pipeline/openai_request.go diff --git a/internal/guardrails/adapter/openai.go b/internal/guardrails/adapter/openai.go new file mode 100644 index 000000000..5858a0d16 --- /dev/null +++ b/internal/guardrails/adapter/openai.go @@ -0,0 +1,318 @@ +package adapter + +import ( + "encoding/json" + "strings" + + "github.com/openai/openai-go/v3" + "github.com/openai/openai-go/v3/packages/param" + "github.com/openai/openai-go/v3/responses" + + guardrailscore "github.com/tingly-dev/tingly-box/internal/guardrails/core" +) + +func AdaptMessagesFromOpenAIChat(messages []openai.ChatCompletionMessageParamUnion) []guardrailscore.Message { + out := make([]guardrailscore.Message, 0, len(messages)) + for _, msg := range messages { + m := openAIParamAsMap(msg) + role, _ := m["role"].(string) + content := textFromOpenAIValue(m["content"]) + if role == "assistant" { + content = strings.TrimSpace(strings.Join(nonEmpty(content, commandTextFromOpenAIChatMessage(m)), "\n")) + } + if role == "" && content == "" { + continue + } + out = append(out, guardrailscore.Message{Role: role, Content: content}) + } + return out +} + +func RefreshInputFromOpenAIChatRequest(input guardrailscore.Input) guardrailscore.Input { + req, _ := input.Payload.Request.(*openai.ChatCompletionNewParams) + if req == nil { + return input + } + text, blockCount, partCount := ExtractOpenAIChatToolResultText(req.Messages) + input.Direction = guardrailscore.DirectionRequest + input.Content = guardrailscore.Content{ + Text: text, + Messages: AdaptMessagesFromOpenAIChat(req.Messages), + } + input.HasToolResult = blockCount > 0 + input.ToolResultBlockCount = blockCount + input.ToolResultPartCount = partCount + if input.Payload.Protocol == "" { + input.Payload.Protocol = "openai_chat" + } + input.Payload.Request = req + return input +} + +func RefreshInputFromOpenAIChatResponse(input guardrailscore.Input, resp *openai.ChatCompletion) guardrailscore.Input { + input.Direction = guardrailscore.DirectionResponse + input.Content = guardrailscore.Content{Messages: input.Content.Messages} + if resp == nil || len(resp.Choices) == 0 { + return input + } + msg := resp.Choices[0].Message + input.Content.Text = strings.TrimSpace(strings.Join(nonEmpty(msg.Content, msg.Refusal), "\n")) + if len(msg.ToolCalls) > 0 { + tc := msg.ToolCalls[0] + switch v := tc.AsAny().(type) { + case openai.ChatCompletionMessageFunctionToolCall: + input.Content.Command = BuildCommandFromRawArguments(v.Function.Name, v.Function.Arguments) + case openai.ChatCompletionMessageCustomToolCall: + input.Content.Command = BuildCommandFromRawArguments(v.Custom.Name, v.Custom.Input) + } + } else if msg.FunctionCall.Name != "" || msg.FunctionCall.Arguments != "" { + input.Content.Command = BuildCommandFromRawArguments(msg.FunctionCall.Name, msg.FunctionCall.Arguments) + } + return input +} + +func ExtractOpenAIChatToolResultText(messages []openai.ChatCompletionMessageParamUnion) (string, int, int) { + for i := len(messages) - 1; i >= 0; i-- { + m := openAIParamAsMap(messages[i]) + role, _ := m["role"].(string) + if role != "tool" && role != "function" { + continue + } + text := textFromOpenAIValue(m["content"]) + if text == "" { + if raw, err := json.Marshal(m); err == nil { + text = string(raw) + } + } + return text, 1, countOpenAIContentParts(m["content"]) + } + return "", 0, 0 +} + +func AdaptMessagesFromOpenAIResponses(req *responses.ResponseNewParams) []guardrailscore.Message { + if req == nil { + return nil + } + if !param.IsOmitted(req.Input.OfString) { + return []guardrailscore.Message{{Role: "user", Content: req.Input.OfString.Value}} + } + items := req.Input.OfInputItemList + out := make([]guardrailscore.Message, 0, len(items)) + for _, item := range items { + m := openAIParamAsMap(item) + role, _ := m["role"].(string) + itemType, _ := m["type"].(string) + if role == "" { + switch itemType { + case "function_call", "custom_tool_call", "mcp_call": + role = "assistant" + case "function_call_output", "custom_tool_call_output", "mcp_call_output": + role = "tool" + } + } + content := textFromOpenAIValue(m["content"]) + if content == "" { + content = textFromOpenAIValue(m["output"]) + } + if role == "assistant" { + content = strings.TrimSpace(strings.Join(nonEmpty(content, commandTextFromOpenAIResponsesItem(m)), "\n")) + } + if role == "" && content == "" { + continue + } + out = append(out, guardrailscore.Message{Role: role, Content: content}) + } + return out +} + +func RefreshInputFromOpenAIResponsesRequest(input guardrailscore.Input) guardrailscore.Input { + req, _ := input.Payload.Request.(*responses.ResponseNewParams) + if req == nil { + return input + } + text, blockCount, partCount := ExtractOpenAIResponsesToolResultText(req) + input.Direction = guardrailscore.DirectionRequest + input.Content = guardrailscore.Content{ + Text: text, + Messages: AdaptMessagesFromOpenAIResponses(req), + } + input.HasToolResult = blockCount > 0 + input.ToolResultBlockCount = blockCount + input.ToolResultPartCount = partCount + if input.Payload.Protocol == "" { + input.Payload.Protocol = "openai_responses" + } + input.Payload.Request = req + return input +} + +func RefreshInputFromOpenAIResponsesResponse(input guardrailscore.Input, resp *responses.Response) guardrailscore.Input { + input.Direction = guardrailscore.DirectionResponse + input.Content = guardrailscore.Content{Messages: input.Content.Messages} + if resp == nil { + return input + } + var textParts []string + for _, item := range resp.Output { + m := openAIParamAsMap(item) + switch m["type"] { + case "message": + textParts = append(textParts, textFromOpenAIValue(m["content"])) + case "function_call", "custom_tool_call", "mcp_call": + if input.Content.Command == nil { + input.Content.Command = commandFromOpenAIResponsesMap(m) + } + } + } + input.Content.Text = strings.TrimSpace(strings.Join(nonEmpty(textParts...), "\n")) + return input +} + +func ExtractOpenAIResponsesToolResultText(req *responses.ResponseNewParams) (string, int, int) { + if req == nil || param.IsOmitted(req.Input.OfInputItemList) { + return "", 0, 0 + } + items := req.Input.OfInputItemList + for i := len(items) - 1; i >= 0; i-- { + m := openAIParamAsMap(items[i]) + itemType, _ := m["type"].(string) + if itemType != "function_call_output" && itemType != "custom_tool_call_output" && itemType != "mcp_call_output" { + continue + } + text := textFromOpenAIValue(m["output"]) + if text == "" { + if raw, err := json.Marshal(m); err == nil { + text = string(raw) + } + } + return text, 1, countOpenAIContentParts(m["output"]) + } + return "", 0, 0 +} + +func commandFromOpenAIResponsesMap(m map[string]interface{}) *guardrailscore.Command { + name, _ := m["name"].(string) + args := "" + if raw, ok := m["arguments"].(string); ok { + args = raw + } else if raw, ok := m["input"].(string); ok { + args = raw + } + return BuildCommandFromRawArguments(name, args) +} + +func commandTextFromOpenAIChatMessage(m map[string]interface{}) string { + if fc, ok := m["function_call"].(map[string]interface{}); ok { + name, _ := fc["name"].(string) + args, _ := fc["arguments"].(string) + return commandPreview(name, args) + } + calls, ok := m["tool_calls"].([]interface{}) + if !ok || len(calls) == 0 { + return "" + } + call, _ := calls[0].(map[string]interface{}) + if fn, ok := call["function"].(map[string]interface{}); ok { + name, _ := fn["name"].(string) + args, _ := fn["arguments"].(string) + return commandPreview(name, args) + } + if custom, ok := call["custom"].(map[string]interface{}); ok { + name, _ := custom["name"].(string) + args, _ := custom["input"].(string) + return commandPreview(name, args) + } + return "" +} + +func commandTextFromOpenAIResponsesItem(m map[string]interface{}) string { + cmd := commandFromOpenAIResponsesMap(m) + if cmd == nil { + return "" + } + raw := "" + if len(cmd.Arguments) > 0 { + if b, err := json.Marshal(cmd.Arguments); err == nil { + raw = string(b) + } + } + return commandPreview(cmd.Name, raw) +} + +func commandPreview(name, args string) string { + if name == "" && args == "" { + return "" + } + if args == "" { + return "command: " + name + } + return "command: " + name + " arguments: " + args +} + +func openAIParamAsMap(v interface{}) map[string]interface{} { + if v == nil { + return nil + } + raw, err := json.Marshal(v) + if err != nil { + return nil + } + var out map[string]interface{} + if err := json.Unmarshal(raw, &out); err != nil { + return nil + } + return out +} + +func textFromOpenAIValue(v interface{}) string { + switch value := v.(type) { + case nil: + return "" + case string: + return value + case []interface{}: + parts := make([]string, 0, len(value)) + for _, item := range value { + if m, ok := item.(map[string]interface{}); ok { + for _, key := range []string{"text", "refusal", "output_text"} { + if text, ok := m[key].(string); ok && text != "" { + parts = append(parts, text) + break + } + } + } else if text := textFromOpenAIValue(item); text != "" { + parts = append(parts, text) + } + } + return strings.Join(parts, "\n") + case map[string]interface{}: + for _, key := range []string{"text", "refusal", "output", "content"} { + if text := textFromOpenAIValue(value[key]); text != "" { + return text + } + } + } + return "" +} + +func countOpenAIContentParts(v interface{}) int { + if parts, ok := v.([]interface{}); ok { + if len(parts) > 0 { + return len(parts) + } + } + if v != nil { + return 1 + } + return 0 +} + +func nonEmpty(values ...string) []string { + out := make([]string, 0, len(values)) + for _, value := range values { + if strings.TrimSpace(value) != "" { + out = append(out, value) + } + } + return out +} diff --git a/internal/guardrails/mutate/openai.go b/internal/guardrails/mutate/openai.go new file mode 100644 index 000000000..b1b3cf22d --- /dev/null +++ b/internal/guardrails/mutate/openai.go @@ -0,0 +1,371 @@ +package mutate + +import ( + "encoding/json" + + "github.com/openai/openai-go/v3" + "github.com/openai/openai-go/v3/packages/param" + "github.com/openai/openai-go/v3/responses" + + guardrailscore "github.com/tingly-dev/tingly-box/internal/guardrails/core" + guardrailsevaluate "github.com/tingly-dev/tingly-box/internal/guardrails/evaluate" +) + +func MaskOpenAIChatRequestCredentials( + req *openai.ChatCompletionNewParams, + credentials []guardrailscore.ProtectedCredential, + state *guardrailscore.CredentialMaskState, +) (bool, bool) { + if req == nil || len(credentials) == 0 { + return false, false + } + messages, ok := openAIChatMessagesToMaps(req.Messages) + if !ok { + return false, false + } + changed := false + latestTurnChanged := false + for i := range messages { + if maskOpenAIMessageMap(messages[i], credentials, state) { + changed = true + if i == len(messages)-1 { + latestTurnChanged = true + } + } + } + if changed { + _ = openAIChatMessagesFromMaps(messages, &req.Messages) + } + return changed, latestTurnChanged +} + +func MaskOpenAIResponsesRequestCredentials( + req *responses.ResponseNewParams, + credentials []guardrailscore.ProtectedCredential, + state *guardrailscore.CredentialMaskState, +) (bool, bool) { + if req == nil || len(credentials) == 0 { + return false, false + } + if !param.IsOmitted(req.Input.OfString) { + if next, ok := guardrailscore.AliasText(req.Input.OfString.Value, credentials, state); ok { + req.Input.OfString = param.NewOpt(next) + return true, true + } + return false, false + } + if param.IsOmitted(req.Input.OfInputItemList) { + return false, false + } + items, ok := openAIResponsesInputItemsToMaps(req.Input.OfInputItemList) + if !ok { + return false, false + } + changed := false + latestTurnChanged := false + for i := range items { + if maskOpenAIMessageMap(items[i], credentials, state) { + changed = true + if i == len(items)-1 { + latestTurnChanged = true + } + } + } + if changed { + _ = openAIResponsesInputItemsFromMaps(items, &req.Input.OfInputItemList) + } + return changed, latestTurnChanged +} + +func MutateOpenAIChatToolResultRequest(req *openai.ChatCompletionNewParams, evaluation guardrailsevaluate.Evaluation) (bool, string) { + if req == nil || evaluation.Result.Verdict != guardrailscore.VerdictBlock { + return false, "" + } + message := BlockMessageForToolResult(evaluation.Result) + if message == "" { + return false, "" + } + messages, ok := openAIChatMessagesToMaps(req.Messages) + if !ok { + return false, "" + } + changed := false + for i := range messages { + role, _ := messages[i]["role"].(string) + if role != "tool" && role != "function" { + continue + } + messages[i]["content"] = message + changed = true + } + if changed { + _ = openAIChatMessagesFromMaps(messages, &req.Messages) + } + return changed, message +} + +func MutateOpenAIResponsesToolResultRequest(req *responses.ResponseNewParams, evaluation guardrailsevaluate.Evaluation) (bool, string) { + if req == nil || evaluation.Result.Verdict != guardrailscore.VerdictBlock || param.IsOmitted(req.Input.OfInputItemList) { + return false, "" + } + message := BlockMessageForToolResult(evaluation.Result) + if message == "" { + return false, "" + } + items, ok := openAIResponsesInputItemsToMaps(req.Input.OfInputItemList) + if !ok { + return false, "" + } + changed := false + for i := range items { + itemType, _ := items[i]["type"].(string) + if itemType != "function_call_output" && itemType != "custom_tool_call_output" && itemType != "mcp_call_output" { + continue + } + items[i]["output"] = message + changed = true + } + if changed { + _ = openAIResponsesInputItemsFromMaps(items, &req.Input.OfInputItemList) + } + return changed, message +} + +func MutateOpenAIChatResponse(resp *openai.ChatCompletion, evaluation guardrailsevaluate.Evaluation) (bool, string) { + if resp == nil || evaluation.Result.Verdict != guardrailscore.VerdictBlock { + return false, "" + } + blockMessage := BlockMessageForEvaluation(evaluation) + if len(resp.Choices) == 0 { + resp.Choices = []openai.ChatCompletionChoice{{Index: 0}} + } + choice := &resp.Choices[0] + choice.Message.Content = blockMessage + choice.Message.Refusal = "" + choice.Message.ToolCalls = nil + choice.Message.FunctionCall = openai.ChatCompletionMessageFunctionCall{} + choice.FinishReason = "stop" + return true, blockMessage +} + +func MutateOpenAIResponsesResponse(resp *responses.Response, evaluation guardrailsevaluate.Evaluation) (bool, string) { + if resp == nil || evaluation.Result.Verdict != guardrailscore.VerdictBlock { + return false, "" + } + blockMessage := BlockMessageForEvaluation(evaluation) + itemID := resp.ID + "_guardrails" + raw := []map[string]interface{}{ + { + "id": itemID, + "type": "message", + "role": "assistant", + "status": "completed", + "content": []map[string]interface{}{ + { + "type": "output_text", + "text": blockMessage, + }, + }, + }, + } + payload, err := json.Marshal(raw) + if err != nil { + return false, "" + } + var output []responses.ResponseOutputItemUnion + if err := json.Unmarshal(payload, &output); err != nil { + return false, "" + } + resp.Output = output + resp.Status = "completed" + return true, blockMessage +} + +func RestoreOpenAIChatResponseCredentials(state *guardrailscore.CredentialMaskState, resp *openai.ChatCompletion) bool { + if state == nil || resp == nil || len(state.AliasToReal) == 0 { + return false + } + changed := false + for i := range resp.Choices { + msg := &resp.Choices[i].Message + if next, ok := guardrailscore.RestoreText(msg.Content, state); ok { + msg.Content = next + changed = true + } + if next, ok := guardrailscore.RestoreText(msg.Refusal, state); ok { + msg.Refusal = next + changed = true + } + for j := range msg.ToolCalls { + if restoreOpenAIChatToolCall(&msg.ToolCalls[j], state) { + changed = true + } + } + } + return changed +} + +func RestoreOpenAIResponsesResponseCredentials(state *guardrailscore.CredentialMaskState, resp *responses.Response) bool { + if state == nil || resp == nil || len(state.AliasToReal) == 0 { + return false + } + raw, err := json.Marshal(resp.Output) + if err != nil || !guardrailscore.MayContainAliasToken(string(raw)) { + return false + } + var parsed interface{} + if err := json.Unmarshal(raw, &parsed); err != nil { + return false + } + restored, changed := guardrailscore.RestoreStructuredValue(parsed, state) + if !changed { + return false + } + payload, err := json.Marshal(restored) + if err != nil { + return false + } + var output []responses.ResponseOutputItemUnion + if err := json.Unmarshal(payload, &output); err != nil { + return false + } + resp.Output = output + return true +} + +func maskOpenAIMessageMap(m map[string]interface{}, credentials []guardrailscore.ProtectedCredential, state *guardrailscore.CredentialMaskState) bool { + changed := false + for _, key := range []string{"content", "output"} { + if next, ok := aliasOpenAIValue(m[key], credentials, state); ok { + m[key] = next + changed = true + } + } + for _, key := range []string{"arguments", "input"} { + if next, ok := aliasOpenAIJSONishString(m[key], credentials, state); ok { + m[key] = next + changed = true + } + } + if calls, ok := m["tool_calls"].([]interface{}); ok { + for _, call := range calls { + callMap, _ := call.(map[string]interface{}) + for _, key := range []string{"function", "custom"} { + child, _ := callMap[key].(map[string]interface{}) + for _, argKey := range []string{"arguments", "input"} { + if next, ok := aliasOpenAIJSONishString(child[argKey], credentials, state); ok { + child[argKey] = next + changed = true + } + } + } + } + } + return changed +} + +func aliasOpenAIValue(value interface{}, credentials []guardrailscore.ProtectedCredential, state *guardrailscore.CredentialMaskState) (interface{}, bool) { + switch v := value.(type) { + case string: + return guardrailscore.AliasText(v, credentials, state) + case []interface{}, map[string]interface{}: + return guardrailscore.AliasStructuredValue(v, credentials, state) + default: + return nil, false + } +} + +func aliasOpenAIJSONishString(value interface{}, credentials []guardrailscore.ProtectedCredential, state *guardrailscore.CredentialMaskState) (interface{}, bool) { + raw, ok := value.(string) + if !ok || raw == "" { + return nil, false + } + var parsed interface{} + if err := json.Unmarshal([]byte(raw), &parsed); err == nil { + if next, changed := guardrailscore.AliasStructuredValue(parsed, credentials, state); changed { + payload, err := json.Marshal(next) + if err == nil { + return string(payload), true + } + } + } + return guardrailscore.AliasText(raw, credentials, state) +} + +func restoreOpenAIChatToolCall(call *openai.ChatCompletionMessageToolCallUnion, state *guardrailscore.CredentialMaskState) bool { + if call == nil { + return false + } + raw := call.RawJSON() + if raw == "" || !guardrailscore.MayContainAliasToken(raw) { + return false + } + var parsed interface{} + if err := json.Unmarshal([]byte(raw), &parsed); err != nil { + return false + } + restored, changed := guardrailscore.RestoreStructuredValue(parsed, state) + if !changed { + return false + } + payload, err := json.Marshal(restored) + if err != nil { + return false + } + var next openai.ChatCompletionMessageToolCallUnion + if err := json.Unmarshal(payload, &next); err != nil { + return false + } + *call = next + return true +} + +func openAIChatMessagesToMaps(messages []openai.ChatCompletionMessageParamUnion) ([]map[string]interface{}, bool) { + raw, err := json.Marshal(messages) + if err != nil { + return nil, false + } + var out []map[string]interface{} + if err := json.Unmarshal(raw, &out); err != nil { + return nil, false + } + return out, true +} + +func openAIChatMessagesFromMaps(messages []map[string]interface{}, target *[]openai.ChatCompletionMessageParamUnion) bool { + raw, err := json.Marshal(messages) + if err != nil { + return false + } + var out []openai.ChatCompletionMessageParamUnion + if err := json.Unmarshal(raw, &out); err != nil { + return false + } + *target = out + return true +} + +func openAIResponsesInputItemsToMaps(items []responses.ResponseInputItemUnionParam) ([]map[string]interface{}, bool) { + raw, err := json.Marshal(items) + if err != nil { + return nil, false + } + var out []map[string]interface{} + if err := json.Unmarshal(raw, &out); err != nil { + return nil, false + } + return out, true +} + +func openAIResponsesInputItemsFromMaps(items []map[string]interface{}, target *responses.ResponseInputParam) bool { + raw, err := json.Marshal(items) + if err != nil { + return false + } + var out responses.ResponseInputParam + if err := json.Unmarshal(raw, &out); err != nil { + return false + } + *target = out + return true +} diff --git a/internal/guardrails/pipeline/openai_nonstream_response.go b/internal/guardrails/pipeline/openai_nonstream_response.go new file mode 100644 index 000000000..f16c1a10c --- /dev/null +++ b/internal/guardrails/pipeline/openai_nonstream_response.go @@ -0,0 +1,66 @@ +package pipeline + +import ( + "context" + + "github.com/openai/openai-go/v3" + "github.com/openai/openai-go/v3/responses" + + guardrails "github.com/tingly-dev/tingly-box/internal/guardrails" + guardrailsadapter "github.com/tingly-dev/tingly-box/internal/guardrails/adapter" + guardrailscore "github.com/tingly-dev/tingly-box/internal/guardrails/core" + guardrailsevaluate "github.com/tingly-dev/tingly-box/internal/guardrails/evaluate" + guardrailsmutate "github.com/tingly-dev/tingly-box/internal/guardrails/mutate" +) + +func ProcessOpenAIChatNonStreamResponse( + ctx context.Context, + runtime *guardrails.Guardrails, + input guardrailscore.Input, + resp *openai.ChatCompletion, +) (NonStreamResponseMutation, error) { + adaptedInput := guardrailsadapter.RefreshInputFromOpenAIChatResponse(input, resp) + evaluation, err := guardrailsevaluate.EvaluateInput(ctx, runtime, adaptedInput) + if err != nil { + adaptedInput.SetContextValue("guardrails_error", err.Error()) + return NonStreamResponseMutation{Input: adaptedInput}, err + } + evaluation.Input.SetContextValue("guardrails_result", evaluation.Result) + changed, blockMessage := guardrailsmutate.MutateOpenAIChatResponse(resp, evaluation) + if changed { + evaluation.Input.SetContextValue("guardrails_block_message", blockMessage) + runtime.AddHistory(evaluation.Input, evaluation.Result, "response", blockMessage) + } + return NonStreamResponseMutation{ + Input: evaluation.Input, + Evaluation: evaluation, + Changed: changed, + BlockMessage: blockMessage, + }, nil +} + +func ProcessOpenAIResponsesNonStreamResponse( + ctx context.Context, + runtime *guardrails.Guardrails, + input guardrailscore.Input, + resp *responses.Response, +) (NonStreamResponseMutation, error) { + adaptedInput := guardrailsadapter.RefreshInputFromOpenAIResponsesResponse(input, resp) + evaluation, err := guardrailsevaluate.EvaluateInput(ctx, runtime, adaptedInput) + if err != nil { + adaptedInput.SetContextValue("guardrails_error", err.Error()) + return NonStreamResponseMutation{Input: adaptedInput}, err + } + evaluation.Input.SetContextValue("guardrails_result", evaluation.Result) + changed, blockMessage := guardrailsmutate.MutateOpenAIResponsesResponse(resp, evaluation) + if changed { + evaluation.Input.SetContextValue("guardrails_block_message", blockMessage) + runtime.AddHistory(evaluation.Input, evaluation.Result, "response", blockMessage) + } + return NonStreamResponseMutation{ + Input: evaluation.Input, + Evaluation: evaluation, + Changed: changed, + BlockMessage: blockMessage, + }, nil +} diff --git a/internal/guardrails/pipeline/openai_request.go b/internal/guardrails/pipeline/openai_request.go new file mode 100644 index 000000000..9caab0362 --- /dev/null +++ b/internal/guardrails/pipeline/openai_request.go @@ -0,0 +1,93 @@ +package pipeline + +import ( + "context" + "strings" + + "github.com/openai/openai-go/v3" + "github.com/openai/openai-go/v3/responses" + + guardrails "github.com/tingly-dev/tingly-box/internal/guardrails" + guardrailsadapter "github.com/tingly-dev/tingly-box/internal/guardrails/adapter" + guardrailscore "github.com/tingly-dev/tingly-box/internal/guardrails/core" + guardrailsevaluate "github.com/tingly-dev/tingly-box/internal/guardrails/evaluate" + guardrailsmutate "github.com/tingly-dev/tingly-box/internal/guardrails/mutate" +) + +func ProcessOpenAIChatRequest(ctx context.Context, runtime *guardrails.Guardrails, input guardrailscore.Input) error { + req, ok := input.Payload.Request.(*openai.ChatCompletionNewParams) + if !ok || req == nil { + return nil + } + return processAnthropicRequest( + ctx, + runtime, + input, + "openai_chat", + guardrailsadapter.RefreshInputFromOpenAIChatRequest, + EvaluateOpenAIChatToolResultRequest, + func(credentials []guardrailscore.ProtectedCredential, state *guardrailscore.CredentialMaskState) (bool, bool) { + return guardrailsmutate.MaskOpenAIChatRequestCredentials(req, credentials, state) + }, + ) +} + +func ProcessOpenAIResponsesRequest(ctx context.Context, runtime *guardrails.Guardrails, input guardrailscore.Input) error { + req, ok := input.Payload.Request.(*responses.ResponseNewParams) + if !ok || req == nil { + return nil + } + return processAnthropicRequest( + ctx, + runtime, + input, + "openai_responses", + guardrailsadapter.RefreshInputFromOpenAIResponsesRequest, + EvaluateOpenAIResponsesToolResultRequest, + func(credentials []guardrailscore.ProtectedCredential, state *guardrailscore.CredentialMaskState) (bool, bool) { + return guardrailsmutate.MaskOpenAIResponsesRequestCredentials(req, credentials, state) + }, + ) +} + +func EvaluateOpenAIChatToolResultRequest(ctx context.Context, runtime *guardrails.Guardrails, input guardrailscore.Input) (ToolResultMutation, error) { + req, _ := input.Payload.Request.(*openai.ChatCompletionNewParams) + if req == nil || !input.HasToolResult { + return ToolResultMutation{Input: input}, nil + } + if strings.HasPrefix(input.Content.Text, guardrailsadapter.BlockPrefix) { + return ToolResultMutation{Input: input}, nil + } + evaluation, err := guardrailsevaluate.EvaluateInput(ctx, runtime, input) + if err != nil { + return ToolResultMutation{}, err + } + changed, message := guardrailsmutate.MutateOpenAIChatToolResultRequest(req, evaluation) + return ToolResultMutation{ + Input: input, + Evaluation: evaluation, + Changed: changed, + Message: message, + }, nil +} + +func EvaluateOpenAIResponsesToolResultRequest(ctx context.Context, runtime *guardrails.Guardrails, input guardrailscore.Input) (ToolResultMutation, error) { + req, _ := input.Payload.Request.(*responses.ResponseNewParams) + if req == nil || !input.HasToolResult { + return ToolResultMutation{Input: input}, nil + } + if strings.HasPrefix(input.Content.Text, guardrailsadapter.BlockPrefix) { + return ToolResultMutation{Input: input}, nil + } + evaluation, err := guardrailsevaluate.EvaluateInput(ctx, runtime, input) + if err != nil { + return ToolResultMutation{}, err + } + changed, message := guardrailsmutate.MutateOpenAIResponsesToolResultRequest(req, evaluation) + return ToolResultMutation{ + Input: input, + Evaluation: evaluation, + Changed: changed, + Message: message, + }, nil +} diff --git a/internal/server/guardrails_runtime.go b/internal/server/guardrails_runtime.go index 768456b6b..4d011defe 100644 --- a/internal/server/guardrails_runtime.go +++ b/internal/server/guardrails_runtime.go @@ -3,6 +3,8 @@ package server import ( "github.com/anthropics/anthropic-sdk-go" "github.com/gin-gonic/gin" + "github.com/openai/openai-go/v3" + "github.com/openai/openai-go/v3/responses" "github.com/sirupsen/logrus" "github.com/tingly-dev/tingly-box/internal/guardrails" @@ -15,6 +17,8 @@ import ( ) var guardrailsSupportedScenarios = []string{ + string(typ.ScenarioOpenAI), + string(typ.ScenarioCodex), string(typ.ScenarioAnthropic), string(typ.ScenarioClaudeCode), } @@ -352,6 +356,46 @@ func (s *Server) applyGuardrailsToAnthropicV1BetaRequest(c *gin.Context, req *an } } +func (s *Server) applyGuardrailsToOpenAIChatRequest(c *gin.Context, req *openai.ChatCompletionNewParams, actualModel string, provider *typ.Provider) { + if req == nil { + return + } + + input := s.buildGuardrailsBaseInput(c, actualModel, provider, guardrailscore.DirectionRequest, nil) + input.State.CredentialMask = ensureGuardrailsCredentialMaskState(c) + input.Payload.Protocol = "openai_chat" + input.Payload.Request = req + + err := guardrailspipeline.ProcessOpenAIChatRequest( + c.Request.Context(), + s.currentGuardrailsRuntime(), + input, + ) + if err != nil { + return + } +} + +func (s *Server) applyGuardrailsToOpenAIResponsesRequest(c *gin.Context, req *responses.ResponseNewParams, actualModel string, provider *typ.Provider) { + if req == nil { + return + } + + input := s.buildGuardrailsBaseInput(c, actualModel, provider, guardrailscore.DirectionRequest, nil) + input.State.CredentialMask = ensureGuardrailsCredentialMaskState(c) + input.Payload.Protocol = "openai_responses" + input.Payload.Request = req + + err := guardrailspipeline.ProcessOpenAIResponsesRequest( + c.Request.Context(), + s.currentGuardrailsRuntime(), + input, + ) + if err != nil { + return + } +} + // ---------------------------------------------------------------------- // Non-Stream Response Guardrails // ---------------------------------------------------------------------- @@ -403,3 +447,47 @@ func (s *Server) applyGuardrailsToAnthropicV1BetaNonStreamResponse(c *gin.Contex } return mutation.Changed } + +func (s *Server) applyGuardrailsToOpenAIChatNonStreamResponse(c *gin.Context, req *openai.ChatCompletionNewParams, actualModel string, provider *typ.Provider, resp *openai.ChatCompletion) bool { + if req == nil || resp == nil { + return false + } + + maskState := ensureGuardrailsCredentialMaskState(c) + messageHistory := guardrailsadapter.AdaptMessagesFromOpenAIChat(req.Messages) + input := s.buildGuardrailsBaseInput(c, actualModel, provider, guardrailscore.DirectionResponse, messageHistory) + input.State.CredentialMask = maskState + input.Payload.Protocol = "openai_chat" + input.Payload.Response = resp + + mutation, err := guardrailspipeline.ProcessOpenAIChatNonStreamResponse(c.Request.Context(), s.currentGuardrailsRuntime(), input, resp) + if err != nil { + return false + } + if !mutation.Changed { + guardrailsmutate.RestoreOpenAIChatResponseCredentials(maskState, resp) + } + return mutation.Changed +} + +func (s *Server) applyGuardrailsToOpenAIResponsesNonStreamResponse(c *gin.Context, req *responses.ResponseNewParams, actualModel string, provider *typ.Provider, resp *responses.Response) bool { + if req == nil || resp == nil { + return false + } + + maskState := ensureGuardrailsCredentialMaskState(c) + messageHistory := guardrailsadapter.AdaptMessagesFromOpenAIResponses(req) + input := s.buildGuardrailsBaseInput(c, actualModel, provider, guardrailscore.DirectionResponse, messageHistory) + input.State.CredentialMask = maskState + input.Payload.Protocol = "openai_responses" + input.Payload.Response = resp + + mutation, err := guardrailspipeline.ProcessOpenAIResponsesNonStreamResponse(c.Request.Context(), s.currentGuardrailsRuntime(), input, resp) + if err != nil { + return false + } + if !mutation.Changed { + guardrailsmutate.RestoreOpenAIResponsesResponseCredentials(maskState, resp) + } + return mutation.Changed +} diff --git a/internal/server/mcp_anthropic_v1_helper.go b/internal/server/mcp_anthropic_v1_helper.go index 5e0e23739..a0fd72079 100644 --- a/internal/server/mcp_anthropic_v1_helper.go +++ b/internal/server/mcp_anthropic_v1_helper.go @@ -463,6 +463,11 @@ func (s *Server) dispatchGenericOpenAIChatNonStream( // Update affinity s.updateAffinityMessageID(c, rule, string(response.ID)) + _, _, _, _, scenario, _, _ := GetTrackingContext(c) + if s.guardrailsEnabledForScenario(scenario) { + s.applyGuardrailsToOpenAIChatNonStreamResponse(c, req, reqCtx.RequestModel, provider, response) + } + // Return response (OpenAI format) c.JSON(http.StatusOK, response) } diff --git a/internal/server/openai.go b/internal/server/openai.go index b5f59ff3e..7d7cdcc83 100644 --- a/internal/server/openai.go +++ b/internal/server/openai.go @@ -230,6 +230,11 @@ func (s *Server) OpenAIChatCompletion(c *gin.Context, req protocol.OpenAIChatCom } transform.AlignToolMessagesForOpenAI(&req.ChatCompletionNewParams) + _, _, _, _, scenario, _, _ := GetTrackingContext(c) + if s.guardrailsEnabledForScenario(scenario) { + s.applyGuardrailsToOpenAIChatRequest(c, &req.ChatCompletionNewParams, actualModel, provider) + } + // === Cap max_tokens at model's maximum === if req.MaxTokens.Valid() && req.MaxTokens.Value > int64(maxAllowed) { req.MaxTokens.Value = int64(maxAllowed) diff --git a/internal/server/openai_chat.go b/internal/server/openai_chat.go index 182539e1b..55f791011 100644 --- a/internal/server/openai_chat.go +++ b/internal/server/openai_chat.go @@ -65,6 +65,11 @@ func (s *Server) handleNonStreamingRequest(c *gin.Context, provider *typ.Provide usage := protocol.NewTokenUsageWithCache(inputTokens, outputTokens, cacheTokens) s.trackUsageWithTokenUsage(c, usage, nil) + _, _, _, _, scenario, _, _ := GetTrackingContext(c) + if s.guardrailsEnabledForScenario(scenario) { + s.applyGuardrailsToOpenAIChatNonStreamResponse(c, req, string(req.Model), provider, response) + } + // Convert response to JSON map for modification responseJSON, err := json.Marshal(response) if err != nil { diff --git a/internal/server/openai_responses.go b/internal/server/openai_responses.go index 44d571ec1..ee2cf26f1 100644 --- a/internal/server/openai_responses.go +++ b/internal/server/openai_responses.go @@ -144,6 +144,10 @@ func (s *Server) HandleResponsesCreate(c *gin.Context) { req.ResponseNewParams = params // req.Model is replaced with actualModel (resolved backend model) from this point on req.Model = actualModel + + if s.guardrailsEnabledForScenario(scenario) { + s.applyGuardrailsToOpenAIResponsesRequest(c, &req.ResponseNewParams, actualModel, provider) + } s.ResponsesCreate(c, scenarioType, provider, rule, req, rule.RequestModel, maxAllowed) } diff --git a/internal/server/protocol_dispatch.go b/internal/server/protocol_dispatch.go index 6a33528c7..22efb653b 100644 --- a/internal/server/protocol_dispatch.go +++ b/internal/server/protocol_dispatch.go @@ -909,6 +909,11 @@ func (s *Server) nonstreamResponsesToChat(c *gin.Context, reqCtx *transform.Tran tokenUsage := protocol.NewTokenUsageWithCache(inputTokens, outputTokens, cacheTokens) s.trackUsageWithTokenUsage(c, tokenUsage, nil) + _, _, _, _, scenario, _, _ := GetTrackingContext(c) + if s.guardrailsEnabledForScenario(scenario) { + s.applyGuardrailsToOpenAIResponsesNonStreamResponse(c, req, actualModel, provider, responsesResp) + } + chatResp := nonstream.OpenAIResponsesToChat(responsesResp, responseModel) if recorder != nil { recorder.SetAssembledResponse(chatResp) @@ -953,6 +958,11 @@ func (s *Server) nonstreamOpenAIResponses(c *gin.Context, reqCtx *transform.Tran // Track usage s.trackUsageWithTokenUsage(c, protocol.NewTokenUsageWithCache(int(inputTokens), int(outputTokens), int(cacheTokens)), nil) + _, _, _, _, scenario, _, _ := GetTrackingContext(c) + if s.guardrailsEnabledForScenario(scenario) { + s.applyGuardrailsToOpenAIResponsesNonStreamResponse(c, params, reqCtx.RequestModel, provider, response) + } + // Override model in response if needed if responseModel != reqCtx.RequestModel { // Create a copy of the response with updated model diff --git a/internal/server/server.go b/internal/server/server.go index f5048804a..aac7bb994 100644 --- a/internal/server/server.go +++ b/internal/server/server.go @@ -264,6 +264,9 @@ func (s *Server) guardrailsEnabled() bool { return false } return s.config.GetScenarioFlag(typ.ScenarioGlobal, "guardrails") || + s.config.GetScenarioFlag(typ.ScenarioOpenAI, "guardrails") || + s.config.GetScenarioFlag(typ.ScenarioCodex, "guardrails") || + s.config.GetScenarioFlag(typ.ScenarioAnthropic, "guardrails") || s.config.GetScenarioFlag(typ.ScenarioClaudeCode, "guardrails") } From 66f428efce347a7b0ccee268da4bc55b2ef92497 Mon Sep 17 00:00:00 2001 From: azhou Date: Thu, 7 May 2026 10:26:41 +0800 Subject: [PATCH 2/2] add stream handler to openai guardrails --- internal/guardrails/adapter/stream.go | 291 ++++++++++++---- internal/guardrails/evaluate/registry.go | 4 +- internal/guardrails/evaluate/registry_test.go | 4 +- .../guardrails/mutate/anthropic_stream.go | 4 + .../mutate/openai_responses_stream.go | 325 ++++++++++++++++++ internal/guardrails/pipeline/stream.go | 2 +- internal/protocol/context.go | 23 +- .../protocol/stream/openai_passthrough.go | 26 +- internal/server/protocol_dispatch.go | 7 + 9 files changed, 605 insertions(+), 81 deletions(-) create mode 100644 internal/guardrails/mutate/openai_responses_stream.go diff --git a/internal/guardrails/adapter/stream.go b/internal/guardrails/adapter/stream.go index 6a9c35634..947155d3d 100644 --- a/internal/guardrails/adapter/stream.go +++ b/internal/guardrails/adapter/stream.go @@ -9,6 +9,28 @@ import ( "github.com/openai/openai-go/v3/responses" ) +const ( + streamEventContentBlockDelta = "content_block_delta" + streamEventContentBlockStart = "content_block_start" + streamEventContentBlockStop = "content_block_stop" + + streamEventResponseOutputTextDelta = "response.output_text.delta" + streamEventResponseOutputTextDone = "response.output_text.done" + streamEventResponseFunctionArgsDelta = "response.function_call_arguments.delta" + streamEventResponseFunctionArgsDone = "response.function_call_arguments.done" + streamEventResponseCustomToolInputDelta = "response.custom_tool_call_input.delta" + streamEventResponseCustomToolInputDone = "response.custom_tool_call_input.done" + streamEventResponseMCPArgsDelta = "response.mcp_call_arguments.delta" + streamEventResponseMCPArgsDone = "response.mcp_call_arguments.done" + streamEventResponseOutputItemAdded = "response.output_item.added" + streamEventResponseCompleted = "response.completed" + + streamToolTypeAnthropicToolUse = "tool_use" + streamToolTypeFunctionCall = "function_call" + streamToolTypeCustomToolCall = "custom_tool_call" + streamToolTypeMCPCall = "mcp_call" +) + type StreamToolUse struct { Index int ID string @@ -35,6 +57,8 @@ type streamToolUseState struct { args string } +// Provider entry points. + func (a *StreamAccumulator) IngestAnthropicEvent(evt *anthropic.MessageStreamEventUnion) { if evt == nil { return @@ -105,6 +129,8 @@ func (a *StreamAccumulator) IngestAnyEvent(event interface{}) { } } +// Public state accessors. + func (a *StreamAccumulator) NextBlockIndex() int { if a.hasIndex { return a.lastIndex + 1 @@ -125,6 +151,8 @@ func (a *StreamAccumulator) PopCompletedToolUse() (StreamToolUse, bool) { return state, true } +// Generic JSON stream dispatch. + func (a *StreamAccumulator) ingestRawJSON(raw string) { if raw == "" { return @@ -138,51 +166,68 @@ func (a *StreamAccumulator) ingestRawJSON(raw string) { func (a *StreamAccumulator) ingestEventMap(payload map[string]interface{}) { eventType, _ := payload["type"].(string) - index := a.captureIndex(payload) switch eventType { - case "content_block_delta": + case streamEventContentBlockDelta, streamEventContentBlockStart, streamEventContentBlockStop: + a.ingestAnthropicEventMap(payload) + case streamEventResponseOutputTextDelta, + streamEventResponseOutputTextDone, + streamEventResponseFunctionArgsDelta, + streamEventResponseCustomToolInputDelta, + streamEventResponseMCPArgsDelta, + streamEventResponseFunctionArgsDone, + streamEventResponseCustomToolInputDone, + streamEventResponseMCPArgsDone, + streamEventResponseOutputItemAdded, + streamEventResponseCompleted: + a.ingestOpenAIResponsesEventMap(payload) + } +} + +// Provider-specific dispatch. + +func (a *StreamAccumulator) ingestAnthropicEventMap(payload map[string]interface{}) { + eventType, _ := payload["type"].(string) + index := a.captureAnthropicIndex(payload) + + switch eventType { + case streamEventContentBlockDelta: delta, _ := payload["delta"].(map[string]interface{}) - a.ingestDelta(index, delta) - case "content_block_start": + a.ingestAnthropicDelta(index, delta) + case streamEventContentBlockStart: block, _ := payload["content_block"].(map[string]interface{}) - a.ingestContentBlock(index, block) - case "content_block_stop": - a.ingestContentBlockStop(index) - case "response.output_text.delta": + a.ingestAnthropicContentBlock(index, block) + case streamEventContentBlockStop: + a.completeToolUse(index) + } +} + +func (a *StreamAccumulator) ingestOpenAIResponsesEventMap(payload map[string]interface{}) { + eventType, _ := payload["type"].(string) + + switch eventType { + case streamEventResponseOutputTextDelta: if delta, ok := payload["delta"].(string); ok { a.textBuilder.WriteString(delta) } - case "response.output_text.done": + case streamEventResponseOutputTextDone: if text, ok := payload["text"].(string); ok { a.textBuilder.WriteString(text) } - case "response.function_call_arguments.delta", "response.custom_tool_call_input.delta", "response.mcp_call_arguments.delta": - if delta, ok := payload["delta"].(string); ok { - a.commandArgs.WriteString(delta) - a.commandFound = true - } - case "response.function_call_arguments.done", "response.custom_tool_call_input.done", "response.mcp_call_arguments.done": - if name, ok := payload["name"].(string); ok && name != "" { - a.commandName = name - a.commandFound = true - } - case "response.output_item.added": - item, _ := payload["item"].(map[string]interface{}) - a.ingestOutputItem(item) - case "response.completed": - response, _ := payload["response"].(map[string]interface{}) - if output, ok := response["output"].([]interface{}); ok { - for _, item := range output { - if itemMap, ok := item.(map[string]interface{}); ok { - a.ingestOutputItem(itemMap) - } - } - } + case streamEventResponseFunctionArgsDelta, streamEventResponseCustomToolInputDelta, streamEventResponseMCPArgsDelta: + a.ingestOpenAIResponsesToolArgumentsDelta(payload) + case streamEventResponseFunctionArgsDone, streamEventResponseCustomToolInputDone, streamEventResponseMCPArgsDone: + a.ingestOpenAIResponsesToolArgumentsDone(payload) + case streamEventResponseOutputItemAdded: + a.ingestOpenAIResponsesOutputItemAdded(payload) + case streamEventResponseCompleted: + a.observeOpenAIResponsesCompleted(payload) } } -func (a *StreamAccumulator) ingestDelta(index int, delta map[string]interface{}) { +// Anthropic stream event handling. + +func (a *StreamAccumulator) ingestAnthropicDelta(index int, delta map[string]interface{}) { if delta == nil { return } @@ -201,12 +246,12 @@ func (a *StreamAccumulator) ingestDelta(index int, delta map[string]interface{}) } } -func (a *StreamAccumulator) ingestContentBlock(index int, block map[string]interface{}) { +func (a *StreamAccumulator) ingestAnthropicContentBlock(index int, block map[string]interface{}) { if block == nil { return } blockType, _ := block["type"].(string) - if blockType != "tool_use" && blockType != "function_call" { + if blockType != streamToolTypeAnthropicToolUse && blockType != streamToolTypeFunctionCall { return } if id, ok := block["id"].(string); ok && id != "" { @@ -234,29 +279,14 @@ func (a *StreamAccumulator) ingestContentBlock(index int, block map[string]inter } } -func (a *StreamAccumulator) ingestContentBlockStop(index int) { - if a.toolUses == nil { - return - } - state, ok := a.toolUses[index] - if !ok { - return - } - a.completed = append(a.completed, StreamToolUse{ - Index: state.index, - ID: state.id, - Name: state.name, - Args: state.args, - }) - delete(a.toolUses, index) -} +// OpenAI Responses stream event handling. -func (a *StreamAccumulator) ingestOutputItem(item map[string]interface{}) { +func (a *StreamAccumulator) observeOpenAIResponsesToolItem(item map[string]interface{}) { if item == nil { return } itemType, _ := item["type"].(string) - if itemType != "function_call" && itemType != "custom_tool_call" && itemType != "mcp_call" { + if !isStreamToolItemType(itemType) { return } if id, ok := item["id"].(string); ok && id != "" { @@ -276,29 +306,126 @@ func (a *StreamAccumulator) ingestOutputItem(item map[string]interface{}) { } } -func (a *StreamAccumulator) captureIndex(payload map[string]interface{}) int { +func (a *StreamAccumulator) ingestOpenAIResponsesToolArgumentsDelta(payload map[string]interface{}) { + responseIndex := a.captureOpenAIResponsesOutputIndex(payload) + if delta, ok := payload["delta"].(string); ok { + a.commandArgs.WriteString(delta) + a.commandFound = true + if state := a.getOrCreateToolUse(responseIndex); state != nil { + state.args += delta + } + } +} + +func (a *StreamAccumulator) ingestOpenAIResponsesToolArgumentsDone(payload map[string]interface{}) { + responseIndex := a.captureOpenAIResponsesOutputIndex(payload) + state := a.getOrCreateToolUse(responseIndex) + if state == nil { + return + } + if itemID, ok := payload["item_id"].(string); ok && itemID != "" { + a.lastToolID = itemID + state.id = itemID + } + if name, ok := payload["name"].(string); ok && name != "" { + a.commandName = name + a.commandFound = true + state.name = name + } + if args, ok := payload["arguments"].(string); ok { + state.args = args + } + a.completeToolUse(responseIndex) +} + +func (a *StreamAccumulator) ingestOpenAIResponsesOutputItemAdded(payload map[string]interface{}) { + responseIndex := a.captureOpenAIResponsesOutputIndex(payload) + item, _ := payload["item"].(map[string]interface{}) + a.mergeOpenAIResponsesToolItem(responseIndex, item) + a.observeOpenAIResponsesToolItem(item) +} + +func (a *StreamAccumulator) observeOpenAIResponsesCompleted(payload map[string]interface{}) { + response, _ := payload["response"].(map[string]interface{}) + if output, ok := response["output"].([]interface{}); ok { + for _, item := range output { + if itemMap, ok := item.(map[string]interface{}); ok { + a.observeOpenAIResponsesToolItem(itemMap) + } + } + } +} + +// Provider-specific index extraction. + +func (a *StreamAccumulator) captureAnthropicIndex(payload map[string]interface{}) int { if payload == nil { return 0 } - if raw, ok := payload["index"]; ok { - switch v := raw.(type) { - case float64: - a.lastIndex = int(v) - a.hasIndex = true - return a.lastIndex - case int: - a.lastIndex = v - a.hasIndex = true - return a.lastIndex - case int64: - a.lastIndex = int(v) - a.hasIndex = true - return a.lastIndex - } + if index, ok := numericIndex(payload, "index"); ok { + a.rememberIndex(index) + return index } return 0 } +func (a *StreamAccumulator) captureOpenAIResponsesOutputIndex(payload map[string]interface{}) int { + if payload == nil { + return 0 + } + if index, ok := numericIndex(payload, "output_index"); ok { + a.rememberIndex(index) + return index + } + return a.captureAnthropicIndex(payload) +} + +// Shared tool-use state management. + +func (a *StreamAccumulator) mergeOpenAIResponsesToolItem(index int, item map[string]interface{}) { + if item == nil { + return + } + itemType, _ := item["type"].(string) + if !isStreamToolItemType(itemType) { + return + } + state := a.getOrCreateToolUse(index) + if state == nil { + return + } + if id, ok := item["id"].(string); ok && id != "" { + state.id = id + a.lastToolID = id + } + if name, ok := item["name"].(string); ok && name != "" { + state.name = name + } + if args, ok := item["arguments"].(string); ok && args != "" { + state.args = args + } + if input, ok := item["input"].(string); ok && input != "" { + state.args = input + } +} + +func (a *StreamAccumulator) completeToolUse(index int) { + if a.toolUses == nil { + return + } + state, ok := a.toolUses[index] + if !ok { + return + } + a.completed = append(a.completed, StreamToolUse{ + Index: state.index, + ID: state.id, + Name: state.name, + Args: state.args, + }) + delete(a.toolUses, index) +} + func (a *StreamAccumulator) getOrCreateToolUse(index int) *streamToolUseState { if a.toolUses == nil { a.toolUses = make(map[int]*streamToolUseState) @@ -310,3 +437,31 @@ func (a *StreamAccumulator) getOrCreateToolUse(index int) *streamToolUseState { a.toolUses[index] = state return state } + +func (a *StreamAccumulator) rememberIndex(index int) { + a.lastIndex = index + a.hasIndex = true +} + +func numericIndex(payload map[string]interface{}, key string) (int, bool) { + raw, ok := payload[key] + if !ok { + return 0, false + } + switch v := raw.(type) { + case float64: + return int(v), true + case int: + return v, true + case int64: + return int(v), true + default: + return 0, false + } +} + +func isStreamToolItemType(itemType string) bool { + return itemType == streamToolTypeFunctionCall || + itemType == streamToolTypeCustomToolCall || + itemType == streamToolTypeMCPCall +} diff --git a/internal/guardrails/evaluate/registry.go b/internal/guardrails/evaluate/registry.go index b674d99f4..a61f3d73e 100644 --- a/internal/guardrails/evaluate/registry.go +++ b/internal/guardrails/evaluate/registry.go @@ -6,8 +6,8 @@ import ( guardrailscore "github.com/tingly-dev/tingly-box/internal/guardrails/core" ) -var defaultResourceAccessToolNames = []string{"bash"} -var defaultCommandExecutionToolNames = []string{"bash"} +var defaultResourceAccessToolNames = []string{"bash", "exec_command"} +var defaultCommandExecutionToolNames = []string{"bash", "exec_command"} // Dependencies provides external services needed by some policy kinds. type Dependencies struct { diff --git a/internal/guardrails/evaluate/registry_test.go b/internal/guardrails/evaluate/registry_test.go index 4022bec32..407f576fa 100644 --- a/internal/guardrails/evaluate/registry_test.go +++ b/internal/guardrails/evaluate/registry_test.go @@ -58,7 +58,7 @@ func TestBuildEvaluatorsCreatesResourceAccessPolicyEvaluator(t *testing.T) { if got := policyRule.scope.Content; len(got) != 1 || got[0] != guardrailscore.ContentTypeCommand { t.Fatalf("policyRule.scope.Content = %v", got) } - if got := policyRule.config.ToolNames; len(got) != 1 || got[0] != "bash" { + if got := policyRule.config.ToolNames; len(got) != 2 || got[0] != "bash" || got[1] != "exec_command" { t.Fatalf("policyRule.config.ToolNames = %#v", got) } } @@ -132,7 +132,7 @@ func TestBuildEvaluatorsDefaultsCommandExecutionToolNames(t *testing.T) { if !ok { t.Fatalf("evaluators[0] type = %T, want *OperationPolicy", evaluators[0]) } - if got := policyRule.config.ToolNames; len(got) != 1 || got[0] != "bash" { + if got := policyRule.config.ToolNames; len(got) != 2 || got[0] != "bash" || got[1] != "exec_command" { t.Fatalf("policyRule.config.ToolNames = %#v", got) } } diff --git a/internal/guardrails/mutate/anthropic_stream.go b/internal/guardrails/mutate/anthropic_stream.go index 4090ab244..0e8fdbcd5 100644 --- a/internal/guardrails/mutate/anthropic_stream.go +++ b/internal/guardrails/mutate/anthropic_stream.go @@ -33,6 +33,10 @@ type AnthropicToolUseDecision struct { } func RegisterAnthropicGuardrailsBlock(state *protocol.GuardrailsStreamState, toolID string, index int, message string) { + RegisterGuardrailsBlock(state, toolID, index, message) +} + +func RegisterGuardrailsBlock(state *protocol.GuardrailsStreamState, toolID string, index int, message string) { if state == nil || toolID == "" || message == "" { return } diff --git a/internal/guardrails/mutate/openai_responses_stream.go b/internal/guardrails/mutate/openai_responses_stream.go new file mode 100644 index 000000000..557689b97 --- /dev/null +++ b/internal/guardrails/mutate/openai_responses_stream.go @@ -0,0 +1,325 @@ +package mutate + +import ( + "encoding/json" + + guardrailscore "github.com/tingly-dev/tingly-box/internal/guardrails/core" + "github.com/tingly-dev/tingly-box/internal/protocol" +) + +const ( + openAIResponsesEventOutputItemAdded = "response.output_item.added" + openAIResponsesEventFunctionArgsDelta = "response.function_call_arguments.delta" + openAIResponsesEventFunctionArgsDone = "response.function_call_arguments.done" + openAIResponsesEventOutputItemDone = "response.output_item.done" + openAIResponsesEventOutputTextDelta = "response.output_text.delta" + openAIResponsesEventOutputTextDone = "response.output_text.done" + openAIResponsesEventCompleted = "response.completed" + openAIResponsesItemTypeFunctionCall = "function_call" + openAIResponsesItemTypeMessage = "message" + openAIResponsesContentTypeOutputText = "output_text" + openAIResponsesStatusInProgress = "in_progress" + openAIResponsesStatusCompleted = "completed" + openAIResponsesGuardrailsReplacementSuffix = "_guardrails" +) + +type OpenAIResponsesBufferedEvent = protocol.GuardrailsBufferedEvent + +// RewriteOpenAIResponsesFunctionCallEvent decides whether an OpenAI Responses +// stream event should be buffered, replaced, or flushed. It only rewrites +// function_call-related events and response.completed after a previous block. +func RewriteOpenAIResponsesFunctionCallEvent( + credentialMask *guardrailscore.CredentialMaskState, + streamState *protocol.GuardrailsStreamState, + event map[string]interface{}, +) (bool, []OpenAIResponsesBufferedEvent, error) { + if streamState == nil || event == nil { + return false, nil, nil + } + ensureOpenAIResponsesStreamState(streamState) + + eventType, _ := event["type"].(string) + switch eventType { + case openAIResponsesEventOutputItemAdded: + item, _ := event["item"].(map[string]interface{}) + if itemType, _ := item["type"].(string); itemType != openAIResponsesItemTypeFunctionCall { + return false, nil, nil + } + itemID := stringFromOpenAIEventMap(item, "id") + if itemID == "" { + return false, nil, nil + } + outputIndex := intFromOpenAIEventMap(event, "output_index") + streamState.OpenAIResponsesOutputIDs[outputIndex] = itemID + streamState.OpenAIResponsesToolEvents[itemID] = append( + streamState.OpenAIResponsesToolEvents[itemID], + OpenAIResponsesBufferedEvent{EventType: eventType, Payload: cloneOpenAIEventMap(event)}, + ) + return true, nil, nil + case openAIResponsesEventFunctionArgsDelta, openAIResponsesEventFunctionArgsDone: + itemID := stringFromOpenAIEventMap(event, "item_id") + if itemID == "" { + itemID = streamState.OpenAIResponsesOutputIDs[intFromOpenAIEventMap(event, "output_index")] + } + if itemID == "" { + return false, nil, nil + } + if _, ok := streamState.OpenAIResponsesToolEvents[itemID]; !ok { + return false, nil, nil + } + streamState.OpenAIResponsesToolEvents[itemID] = append( + streamState.OpenAIResponsesToolEvents[itemID], + OpenAIResponsesBufferedEvent{EventType: eventType, Payload: cloneOpenAIEventMap(event)}, + ) + return true, nil, nil + case openAIResponsesEventOutputItemDone: + item, _ := event["item"].(map[string]interface{}) + itemID := stringFromOpenAIEventMap(item, "id") + if itemID == "" { + itemID = streamState.OpenAIResponsesOutputIDs[intFromOpenAIEventMap(event, "output_index")] + } + if itemID == "" { + return false, nil, nil + } + buffered, ok := streamState.OpenAIResponsesToolEvents[itemID] + if !ok { + return false, nil, nil + } + buffered = append(buffered, OpenAIResponsesBufferedEvent{EventType: eventType, Payload: cloneOpenAIEventMap(event)}) + delete(streamState.OpenAIResponsesToolEvents, itemID) + + outputIndex := intFromOpenAIEventMap(event, "output_index") + if blockMessage := consumeOpenAIResponsesBlockMessage(streamState, itemID, outputIndex); blockMessage != "" { + blocked := protocol.GuardrailsOpenAIResponsesBlockedItem{ + ItemID: itemID, + TextItemID: itemID + openAIResponsesGuardrailsReplacementSuffix, + OutputIndex: outputIndex, + Message: blockMessage, + } + streamState.OpenAIResponsesBlocked = append(streamState.OpenAIResponsesBlocked, blocked) + return true, syntheticOpenAIResponsesBlockEvents(event, blocked), nil + } + + rebuilt, ok := RebuildBufferedOpenAIResponsesFunctionCallEvents(credentialMask, buffered) + if ok { + return true, rebuilt, nil + } + return true, buffered, nil + case openAIResponsesEventCompleted: + if len(streamState.OpenAIResponsesBlocked) == 0 { + return false, nil, nil + } + return true, []OpenAIResponsesBufferedEvent{ + { + EventType: eventType, + Payload: rewriteOpenAIResponsesCompletedEvent(event, streamState.OpenAIResponsesBlocked), + }, + }, nil + default: + return false, nil, nil + } +} + +func RebuildBufferedOpenAIResponsesFunctionCallEvents(state *guardrailscore.CredentialMaskState, events []OpenAIResponsesBufferedEvent) ([]OpenAIResponsesBufferedEvent, bool) { + if state == nil || len(state.AliasToReal) == 0 || len(events) == 0 { + return nil, false + } + + changed := false + rebuilt := make([]OpenAIResponsesBufferedEvent, 0, len(events)) + for _, event := range events { + payload := cloneOpenAIEventMap(event.Payload) + switch event.EventType { + case openAIResponsesEventFunctionArgsDelta: + if delta, ok := payload["delta"].(string); ok && guardrailscore.MayContainAliasToken(delta) { + if restored, ok := guardrailscore.RestoreText(delta, state); ok { + payload["delta"] = restored + changed = true + } + } + case openAIResponsesEventFunctionArgsDone: + if args, ok := payload["arguments"].(string); ok && guardrailscore.MayContainAliasToken(args) { + if restored, ok := guardrailscore.RestoreText(args, state); ok { + payload["arguments"] = restored + changed = true + } + } + case openAIResponsesEventOutputItemAdded, openAIResponsesEventOutputItemDone: + item, _ := payload["item"].(map[string]interface{}) + if item != nil { + if args, ok := item["arguments"].(string); ok && guardrailscore.MayContainAliasToken(args) { + if restored, ok := guardrailscore.RestoreText(args, state); ok { + item["arguments"] = restored + changed = true + } + } + } + } + rebuilt = append(rebuilt, OpenAIResponsesBufferedEvent{EventType: event.EventType, Payload: payload}) + } + if !changed { + return nil, false + } + return rebuilt, true +} + +func ensureOpenAIResponsesStreamState(state *protocol.GuardrailsStreamState) { + if state.OpenAIResponsesToolEvents == nil { + state.OpenAIResponsesToolEvents = make(map[string][]protocol.GuardrailsBufferedEvent) + } + if state.OpenAIResponsesOutputIDs == nil { + state.OpenAIResponsesOutputIDs = make(map[int]string) + } +} + +func consumeOpenAIResponsesBlockMessage(state *protocol.GuardrailsStreamState, itemID string, outputIndex int) string { + if state == nil { + return "" + } + if message, ok := state.PendingBlockMessages[itemID]; ok { + delete(state.PendingBlockMessages, itemID) + delete(state.PendingBlockedIndex, outputIndex) + return message + } + if blockedID, ok := state.PendingBlockedIndex[outputIndex]; ok { + if message, ok := state.PendingBlockMessages[blockedID]; ok { + delete(state.PendingBlockMessages, blockedID) + delete(state.PendingBlockedIndex, outputIndex) + return message + } + } + return "" +} + +func syntheticOpenAIResponsesBlockEvents(source map[string]interface{}, blocked protocol.GuardrailsOpenAIResponsesBlockedItem) []OpenAIResponsesBufferedEvent { + seq := source["sequence_number"] + return []OpenAIResponsesBufferedEvent{ + { + EventType: openAIResponsesEventOutputItemAdded, + Payload: map[string]interface{}{ + "type": openAIResponsesEventOutputItemAdded, + "sequence_number": seq, + "output_index": blocked.OutputIndex, + "item": openAIResponsesBlockMessageItem(blocked, openAIResponsesStatusInProgress, ""), + }, + }, + { + EventType: openAIResponsesEventOutputTextDelta, + Payload: map[string]interface{}{ + "type": openAIResponsesEventOutputTextDelta, + "sequence_number": seq, + "item_id": blocked.TextItemID, + "output_index": blocked.OutputIndex, + "content_index": 0, + "delta": blocked.Message, + "logprobs": []interface{}{}, + }, + }, + { + EventType: openAIResponsesEventOutputTextDone, + Payload: map[string]interface{}{ + "type": openAIResponsesEventOutputTextDone, + "sequence_number": seq, + "item_id": blocked.TextItemID, + "output_index": blocked.OutputIndex, + "content_index": 0, + "text": blocked.Message, + "logprobs": []interface{}{}, + }, + }, + { + EventType: openAIResponsesEventOutputItemDone, + Payload: map[string]interface{}{ + "type": openAIResponsesEventOutputItemDone, + "sequence_number": seq, + "output_index": blocked.OutputIndex, + "item": openAIResponsesBlockMessageItem(blocked, openAIResponsesStatusCompleted, blocked.Message), + }, + }, + } +} + +func rewriteOpenAIResponsesCompletedEvent(event map[string]interface{}, blockedItems []protocol.GuardrailsOpenAIResponsesBlockedItem) map[string]interface{} { + next := cloneOpenAIEventMap(event) + response, _ := next["response"].(map[string]interface{}) + if response == nil { + return next + } + blockedIDs := make(map[string]struct{}, len(blockedItems)) + for _, blocked := range blockedItems { + blockedIDs[blocked.ItemID] = struct{}{} + } + output, _ := response["output"].([]interface{}) + rewritten := make([]interface{}, 0, len(output)+len(blockedItems)) + for _, item := range output { + itemMap, _ := item.(map[string]interface{}) + if itemMap != nil { + if itemType, _ := itemMap["type"].(string); itemType == openAIResponsesItemTypeFunctionCall { + if _, ok := blockedIDs[stringFromOpenAIEventMap(itemMap, "id")]; ok { + continue + } + } + } + rewritten = append(rewritten, item) + } + for _, blocked := range blockedItems { + rewritten = append(rewritten, openAIResponsesBlockMessageItem(blocked, openAIResponsesStatusCompleted, blocked.Message)) + } + response["output"] = rewritten + return next +} + +func openAIResponsesBlockMessageItem(blocked protocol.GuardrailsOpenAIResponsesBlockedItem, status string, text string) map[string]interface{} { + return map[string]interface{}{ + "id": blocked.TextItemID, + "type": openAIResponsesItemTypeMessage, + "role": "assistant", + "status": status, + "content": []map[string]interface{}{ + { + "type": openAIResponsesContentTypeOutputText, + "text": text, + "annotations": []interface{}{}, + }, + }, + } +} + +func cloneOpenAIEventMap(event map[string]interface{}) map[string]interface{} { + if event == nil { + return nil + } + raw, err := json.Marshal(event) + if err != nil { + return event + } + var out map[string]interface{} + if err := json.Unmarshal(raw, &out); err != nil { + return event + } + return out +} + +func stringFromOpenAIEventMap(m map[string]interface{}, key string) string { + if m == nil { + return "" + } + value, _ := m[key].(string) + return value +} + +func intFromOpenAIEventMap(m map[string]interface{}, key string) int { + if m == nil { + return 0 + } + switch value := m[key].(type) { + case int: + return value + case int64: + return int(value) + case float64: + return int(value) + default: + return 0 + } +} diff --git a/internal/guardrails/pipeline/stream.go b/internal/guardrails/pipeline/stream.go index 36bfce89b..d850a42c9 100644 --- a/internal/guardrails/pipeline/stream.go +++ b/internal/guardrails/pipeline/stream.go @@ -114,6 +114,6 @@ func handleGuardrailsBlock( runtime.AddHistory(input, guardrailscore.Result{Verdict: guardrailscore.VerdictBlock}, "tool_use", blockMessage) } if streamState != nil { - guardrailsmutate.RegisterAnthropicGuardrailsBlock(streamState, toolID, blockIndex, blockMessage) + guardrailsmutate.RegisterGuardrailsBlock(streamState, toolID, blockIndex, blockMessage) } } diff --git a/internal/protocol/context.go b/internal/protocol/context.go index e6daafb24..ab03a1287 100644 --- a/internal/protocol/context.go +++ b/internal/protocol/context.go @@ -58,6 +58,12 @@ type GuardrailsStreamState struct { AnthropicToolEvents map[int][]GuardrailsBufferedEvent // AnthropicToolIDs links the buffered block index back to the provider tool id. AnthropicToolIDs map[int]string + // OpenAIResponsesToolEvents buffers one function_call item until the item is done. + OpenAIResponsesToolEvents map[string][]GuardrailsBufferedEvent + // OpenAIResponsesOutputIDs links Responses output_index back to the provider item id. + OpenAIResponsesOutputIDs map[int]string + // OpenAIResponsesBlocked stores blocked function_call replacements until response.completed. + OpenAIResponsesBlocked []GuardrailsOpenAIResponsesBlockedItem } func (hc *HandleContext) EnsureGuardrails() *HandleGuardrails { @@ -71,10 +77,12 @@ func (hc *HandleContext) EnsureGuardrailsStream() *GuardrailsStreamState { guardrails := hc.EnsureGuardrails() if guardrails.Stream == nil { guardrails.Stream = &GuardrailsStreamState{ - PendingBlockMessages: make(map[string]string), - PendingBlockedIndex: make(map[int]string), - AnthropicToolEvents: make(map[int][]GuardrailsBufferedEvent), - AnthropicToolIDs: make(map[int]string), + PendingBlockMessages: make(map[string]string), + PendingBlockedIndex: make(map[int]string), + AnthropicToolEvents: make(map[int][]GuardrailsBufferedEvent), + AnthropicToolIDs: make(map[int]string), + OpenAIResponsesToolEvents: make(map[string][]GuardrailsBufferedEvent), + OpenAIResponsesOutputIDs: make(map[int]string), } } return guardrails.Stream @@ -85,6 +93,13 @@ type GuardrailsBufferedEvent struct { Payload map[string]interface{} } +type GuardrailsOpenAIResponsesBlockedItem struct { + ItemID string + TextItemID string + OutputIndex int + Message string +} + // WithOnStreamEvent adds a hook that is called for each stream event. // Multiple hooks can be added and will be called in order. func (hc *HandleContext) WithOnStreamEvent(hook func(interface{}) error) *HandleContext { diff --git a/internal/protocol/stream/openai_passthrough.go b/internal/protocol/stream/openai_passthrough.go index 5dca175c7..75c57ac1f 100644 --- a/internal/protocol/stream/openai_passthrough.go +++ b/internal/protocol/stream/openai_passthrough.go @@ -15,6 +15,7 @@ import ( openaistream "github.com/openai/openai-go/v3/packages/ssestream" "github.com/openai/openai-go/v3/responses" "github.com/sirupsen/logrus" + guardrailsmutate "github.com/tingly-dev/tingly-box/internal/guardrails/mutate" "github.com/tingly-dev/tingly-box/internal/protocol" "github.com/tingly-dev/tingly-box/internal/protocol/token" ) @@ -364,6 +365,11 @@ func HandleOpenAIResponsesStream(hc *protocol.HandleContext, stream *openaistrea } evt := stream.Current() + for _, hook := range hc.OnStreamEventHooks { + if err := hook(&evt); err != nil { + logrus.WithError(err).Warn("guardrails hook error") + } + } // Accumulate usage from completed events if evt.Response.Usage.InputTokens > 0 { @@ -391,15 +397,16 @@ func HandleOpenAIResponsesStream(hc *protocol.HandleContext, stream *openaistrea // Marshal event using RawJSON() to avoid serializing empty union fields jsonBytes := []byte(evt.RawJSON()) + var parsedEvent map[string]interface{} + // Apply model override if the event contains a response object with a model field if len(jsonBytes) > 0 { - var parsed map[string]interface{} - if err := json.Unmarshal(jsonBytes, &parsed); err == nil { + if err := json.Unmarshal(jsonBytes, &parsedEvent); err == nil { // Check if this event has a response field with a model - if response, ok := parsed["response"].(map[string]interface{}); ok { + if response, ok := parsedEvent["response"].(map[string]interface{}); ok { if model, ok2 := response["model"].(string); ok2 && model != "" { response["model"] = responseModel - modified, err := json.Marshal(parsed) + modified, err := json.Marshal(parsedEvent) if err == nil { jsonBytes = modified } @@ -408,6 +415,17 @@ func HandleOpenAIResponsesStream(hc *protocol.HandleContext, stream *openaistrea } } + if hc.Guardrails != nil && hc.Guardrails.Enabled { + if handled, rewritten, err := guardrailsmutate.RewriteOpenAIResponsesFunctionCallEvent(hc.Guardrails.CredentialMask, hc.Guardrails.Stream, parsedEvent); err != nil { + logrus.WithError(err).Warn("OpenAI Responses guardrails rewrite error") + } else if handled { + for _, rewrittenEvent := range rewritten { + OpenAISSE(c, rewrittenEvent.Payload) + } + return true + } + } + // Send SSE event with event type (e.g., "response.created", "response.output_text.delta") OpenAISSE(c, json.RawMessage(jsonBytes)) return true diff --git a/internal/server/protocol_dispatch.go b/internal/server/protocol_dispatch.go index 22efb653b..a9cb55d02 100644 --- a/internal/server/protocol_dispatch.go +++ b/internal/server/protocol_dispatch.go @@ -14,6 +14,7 @@ import ( "github.com/openai/openai-go/v3" "github.com/openai/openai-go/v3/responses" "github.com/sirupsen/logrus" + guardrailsadapter "github.com/tingly-dev/tingly-box/internal/guardrails/adapter" mcpruntime "github.com/tingly-dev/tingly-box/internal/mcp/runtime" "github.com/tingly-dev/tingly-box/internal/protocol" "github.com/tingly-dev/tingly-box/internal/protocol/nonstream" @@ -1020,6 +1021,12 @@ func (s *Server) streamOpenAIResponses(c *gin.Context, reqCtx *transform.Transfo } return nil }) + + _, _, _, _, scenario, _, _ := GetTrackingContext(c) + if s.guardrailsEnabledForScenario(scenario) { + hc.EnsureGuardrails().Enabled = true + s.attachGuardrailsHooks(c, hc, reqCtx.RequestModel, provider, guardrailsadapter.AdaptMessagesFromOpenAIResponses(params)) + } usage, err := stream.HandleOpenAIResponsesStream(hc, respStream, responseModel) // Track usage from stream handler