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

347 lines
8.7 KiB
Go

package openai
import (
"encoding/json"
"fmt"
"sort"
"strings"
)
var (
textToolCallOpeningHint = "<tool_call"
textMustacheToolCallOpenHint = "{{"
)
type xmlTextToolCallBlock struct {
start int
end int
name string
params map[string]string
err error
}
type textToolSpec struct {
name string
parameters map[string]any
}
type textToolSynthesisResult struct {
cleaned string
toolCalls []any
candidateFound bool
validationErr error
}
type candidateMatch struct {
start int
end int
name string
arguments string
validationErr error
}
func collectXMLCandidateMatches(content string, specs map[string]textToolSpec, names map[string]struct{}) ([]candidateMatch, bool) {
blocks, foundAny := scanXMLTextToolCallBlocks(content)
if !foundAny {
return nil, false
}
out := make([]candidateMatch, 0)
for _, block := range blocks {
if block.err != nil {
out = append(out, candidateMatch{
start: block.start,
end: block.end,
validationErr: block.err,
})
continue
}
if _, ok := names[block.name]; !ok {
out = append(out, candidateMatch{
start: block.start,
end: block.end,
name: block.name,
validationErr: fmt.Errorf("tool %q is not in request tools", block.name),
})
continue
}
spec := specs[block.name]
args := map[string]any{}
for key, value := range block.params {
if key == "" {
continue
}
args[key] = parseTextToolParameterValue(value)
}
args = normalizeTextToolCallArguments(block.name, args, spec)
encoded, err := json.Marshal(args)
if err != nil {
out = append(out, candidateMatch{
start: block.start,
end: block.end,
name: block.name,
validationErr: fmt.Errorf("failed to marshal arguments for tool %q: %w", block.name, err),
})
continue
}
out = append(out, candidateMatch{
start: block.start,
end: block.end,
name: block.name,
arguments: string(encoded),
})
}
return out, foundAny
}
func collectMustacheCandidateMatches(content string, specs map[string]textToolSpec, names map[string]struct{}) ([]candidateMatch, bool) {
if !strings.Contains(content, textMustacheToolCallOpenHint) {
return nil, false
}
out := make([]candidateMatch, 0)
foundAny := false
for start := 0; start < len(content); {
open := strings.Index(content[start:], textMustacheToolCallOpenHint)
if open < 0 {
break
}
blockStart := start + open
nameStart := skipASCIISpaces(content, blockStart+len(textMustacheToolCallOpenHint))
nameEnd := nameStart
for nameEnd < len(content) && isTextToolIdentifierChar(content[nameEnd], nameEnd == nameStart) {
nameEnd++
}
name := content[nameStart:nameEnd]
if name == "" {
start = blockStart + len(textMustacheToolCallOpenHint)
continue
}
callOpen := skipASCIISpaces(content, nameEnd)
if callOpen >= len(content) || content[callOpen] != '(' {
start = blockStart + len(textMustacheToolCallOpenHint)
continue
}
foundAny = true
argsStart := callOpen + 1
argsEnd, blockEnd, ok := findMustacheToolCallEnd(content, argsStart)
if !ok {
out = append(out, candidateMatch{
start: blockStart,
end: len(content),
name: name,
validationErr: fmt.Errorf("malformed tool call: unclosed mustache tool call %q", name),
})
break
}
if _, ok := names[name]; !ok {
out = append(out, candidateMatch{
start: blockStart,
end: blockEnd,
name: name,
validationErr: fmt.Errorf("tool %q is not in request tools", name),
})
start = blockEnd
continue
}
spec := specs[name]
args, ok := parseMustacheToolCallArguments(content[argsStart:argsEnd])
if !ok {
out = append(out, candidateMatch{
start: blockStart,
end: blockEnd,
name: name,
validationErr: fmt.Errorf("failed to parse arguments for tool %q", name),
})
start = blockEnd
continue
}
args = normalizeTextToolCallArguments(name, args, spec)
encoded, err := json.Marshal(args)
if err != nil {
out = append(out, candidateMatch{
start: blockStart,
end: blockEnd,
name: name,
validationErr: fmt.Errorf("failed to marshal arguments for tool %q: %w", name, err),
})
start = blockEnd
continue
}
out = append(out, candidateMatch{
start: blockStart,
end: blockEnd,
name: name,
arguments: string(encoded),
})
start = blockEnd
}
return out, foundAny
}
func synthesizeToolCallsFromTextResult(content string, tools []any, runID string) textToolSynthesisResult {
specs := requestedToolSpecs(tools)
names := requestedToolNames(tools)
xmlMatches, xmlFound := collectXMLCandidateMatches(content, specs, names)
mustacheMatches, mustacheFound := collectMustacheCandidateMatches(content, specs, names)
candidateFound := xmlFound || mustacheFound
var matches []candidateMatch
matches = append(matches, xmlMatches...)
matches = append(matches, mustacheMatches...)
sort.SliceStable(matches, func(i, j int) bool {
if matches[i].start == matches[j].start {
return matches[i].end > matches[j].end
}
return matches[i].start < matches[j].start
})
var cleaned strings.Builder
cursor := 0
toolCalls := make([]any, 0)
var firstErr error
for _, match := range matches {
if match.start < cursor {
continue
}
if match.validationErr != nil {
if firstErr == nil {
firstErr = match.validationErr
}
continue
}
cleaned.WriteString(content[cursor:match.start])
cursor = match.end
toolCalls = append(toolCalls, map[string]any{
"id": textToolCallID(runID, len(toolCalls)+1),
"type": "function",
"function": map[string]any{
"name": match.name,
"arguments": match.arguments,
},
})
}
if firstErr != nil {
return textToolSynthesisResult{
cleaned: content,
toolCalls: nil,
candidateFound: candidateFound,
validationErr: firstErr,
}
}
if len(toolCalls) == 0 {
return textToolSynthesisResult{
cleaned: content,
toolCalls: nil,
candidateFound: candidateFound,
validationErr: nil,
}
}
cleaned.WriteString(content[cursor:])
return textToolSynthesisResult{
cleaned: strings.TrimSpace(cleaned.String()),
toolCalls: toolCalls,
candidateFound: candidateFound,
validationErr: nil,
}
}
func synthesizeToolCallsFromText(content string, tools []any, runID string) (string, []any) {
res := synthesizeToolCallsFromTextResult(content, tools, runID)
if res.validationErr != nil {
return content, nil
}
return res.cleaned, res.toolCalls
}
func textToolCallID(runID string, ordinal int) string {
normalized := normalizeToolCallIDPart(runID)
if normalized == "" {
return fmt.Sprintf("call_iop_%d", ordinal)
}
return fmt.Sprintf("call_iop_%s_%d", normalized, ordinal)
}
func normalizeToolCallIDPart(s string) string {
var b strings.Builder
for _, r := range s {
switch {
case r >= 'a' && r <= 'z':
b.WriteRune(r)
case r >= 'A' && r <= 'Z':
b.WriteRune(r)
case r >= '0' && r <= '9':
b.WriteRune(r)
case r == '_' || r == '-':
b.WriteRune(r)
default:
b.WriteByte('_')
}
}
return strings.Trim(b.String(), "_-")
}
func requestedToolNames(tools []any) map[string]struct{} {
names := make(map[string]struct{}, len(tools))
for _, tool := range tools {
m, ok := tool.(map[string]any)
if !ok {
continue
}
if name, ok := m["name"].(string); ok && strings.TrimSpace(name) != "" {
names[strings.TrimSpace(name)] = struct{}{}
}
fn, ok := m["function"].(map[string]any)
if !ok {
continue
}
if name, ok := fn["name"].(string); ok && strings.TrimSpace(name) != "" {
names[strings.TrimSpace(name)] = struct{}{}
}
}
return names
}
func requestedToolSpecs(tools []any) map[string]textToolSpec {
specs := make(map[string]textToolSpec, len(tools))
for _, tool := range tools {
m, ok := tool.(map[string]any)
if !ok {
continue
}
if name, ok := m["name"].(string); ok && strings.TrimSpace(name) != "" {
name = strings.TrimSpace(name)
specs[name] = textToolSpec{name: name, parameters: mapValue(m["parameters"])}
}
fn, ok := m["function"].(map[string]any)
if !ok {
continue
}
if name, ok := fn["name"].(string); ok && strings.TrimSpace(name) != "" {
name = strings.TrimSpace(name)
specs[name] = textToolSpec{name: name, parameters: mapValue(fn["parameters"])}
}
}
return specs
}
func mapValue(value any) map[string]any {
m, _ := value.(map[string]any)
return m
}
func parseTextToolParameterValue(raw string) any {
trimmed := strings.TrimSpace(raw)
if trimmed == "" {
return ""
}
var decoded any
if err := json.Unmarshal([]byte(trimmed), &decoded); err == nil {
return decoded
}
return trimmed
}