iop/apps/edge/internal/openai/responses_decode.go

142 lines
4.4 KiB
Go

package openai
import (
"encoding/json"
"errors"
"fmt"
"strings"
)
func decodeResponsesRequest(dec *json.Decoder, req *responsesRequest) error {
var raw map[string]json.RawMessage
if err := dec.Decode(&raw); err != nil {
return fmt.Errorf("invalid JSON request")
}
for key := range raw {
switch key {
case "model", "input", "instructions", "stream", "background", "metadata", "max_output_tokens", "temperature", "top_p":
default:
return fmt.Errorf("%s is not supported for /v1/responses", key)
}
}
normalized, err := json.Marshal(raw)
if err != nil {
return fmt.Errorf("invalid JSON request")
}
if err := json.Unmarshal(normalized, req); err != nil {
return fmt.Errorf("invalid /v1/responses request format")
}
if req.MaxOutputTokens != nil && *req.MaxOutputTokens <= 0 {
return fmt.Errorf("max_output_tokens must be greater than zero")
}
if req.Temperature != nil && (*req.Temperature < 0 || *req.Temperature > 2) {
return fmt.Errorf("temperature must be between 0 and 2")
}
if req.TopP != nil && (*req.TopP < 0 || *req.TopP > 1) {
return fmt.Errorf("top_p must be between 0 and 1")
}
return nil
}
// decodeResponsesEnvelope leniently extracts only the routing-relevant fields
// (model, metadata, stream, background) from a /v1/responses request body. It
// does not reject unknown fields so the provider tunnel passthrough can forward
// Codex/Responses payloads verbatim; strict field validation stays on the
// normalized non-provider path.
func decodeResponsesEnvelope(rawBody []byte) (responsesEnvelope, error) {
var env responsesEnvelope
if err := json.Unmarshal(rawBody, &env); err != nil {
return responsesEnvelope{}, fmt.Errorf("invalid JSON request")
}
return env, nil
}
func parseResponsesInput(raw json.RawMessage) (string, error) {
if len(raw) == 0 {
return "", fmt.Errorf("input is required")
}
var s string
if err := json.Unmarshal(raw, &s); err != nil {
return "", fmt.Errorf("input must be a string")
}
if strings.TrimSpace(s) == "" {
return "", fmt.Errorf("input is required")
}
return s, nil
}
func buildResponsesPrompt(instructions, input string) string {
instructions = strings.TrimSpace(instructions)
if instructions == "" {
return input
}
return instructions + "\n\n" + input
}
func parseOpenAIMetadata(raw json.RawMessage) (map[string]string, string, error) {
if len(raw) == 0 || string(raw) == "null" {
return make(map[string]string), "", nil
}
var rawMap map[string]json.RawMessage
if err := json.Unmarshal(raw, &rawMap); err != nil {
return nil, "", fmt.Errorf("metadata must be an object")
}
if len(rawMap) > 16 {
return nil, "", fmt.Errorf("metadata must contain at most 16 keys")
}
flat := make(map[string]string, len(rawMap))
var workspace string
for key, rawValue := range rawMap {
if len(key) > 64 {
return nil, "", fmt.Errorf("metadata key %q exceeds 64 characters", key)
}
if key == "source" {
return nil, "", fmt.Errorf("metadata.source is not supported")
}
if key == "workspace" {
workspaceValue, err := metadataStringValue(key, rawValue)
if err != nil {
return nil, "", err
}
workspace = strings.TrimSpace(workspaceValue)
continue
}
value, err := metadataStringValue(key, rawValue)
if err != nil {
return nil, "", err
}
flat[key] = value
}
return flat, workspace, nil
}
// errProviderRequestValidation is a sentinel for client request validation
// failures that occur inside the normalized provider-pool PrepareRun hook
// (strict decode, stream/background rejection, input parsing). These errors
// are reported as HTTP 400 invalid_request_error, distinct from backend
// dispatch errors that return 502.
var errProviderRequestValidation = errors.New("provider_request_validation")
// isValidationError reports whether err originated from a PrepareRun validation
// failure in the normalized provider-pool path.
func isValidationError(err error) bool {
return errors.Is(err, errProviderRequestValidation)
}
// metadataStringValue extracts a single string value from a json.RawMessage
// for the metadata field. It rejects non-string values and enforces the
// 512-character length cap.
func metadataStringValue(key string, raw json.RawMessage) (string, error) {
var value string
if err := json.Unmarshal(raw, &value); err != nil {
return "", fmt.Errorf("metadata.%s must be a string", key)
}
if len(value) > 512 {
return "", fmt.Errorf("metadata.%s exceeds 512 characters", key)
}
return value, nil
}