Files
sub2api/backend/internal/service/openai_gateway_grok_tool_protocol.go
T
李建琦 6d655c9903
Release / update-version (push) Has been cancelled
Release / build-frontend (push) Has been cancelled
Release / release (push) Has been cancelled
Release / sync-version-file (push) Has been cancelled
CI / shell (push) Canceled after 0s
CI / test (push) Canceled after 0s
CI / frontend (push) Canceled after 0s
CI / golangci-lint (push) Canceled after 0s
Security Scan / backend-security (push) Canceled after 0s
Security Scan / frontend-security (push) Canceled after 0s
Sub2API v1.0 - AI API 网关(二开初始版本,基于上游 Wei-Shaw/sub2api)
2026-08-21 18:30:13 +08:00

285 lines
7.8 KiB
Go

package service
import (
"bufio"
"bytes"
"encoding/json"
"fmt"
"io"
"strings"
"github.com/Wei-Shaw/sub2api/internal/pkg/apicompat"
"github.com/gin-gonic/gin"
"github.com/tidwall/gjson"
)
const grokResponsesClientToolMappingContextKey = "grok_responses_client_tool_mapping"
func adaptResponsesClientToolsForFunctionUpstream(body []byte, upstream string) ([]byte, apicompat.ResponsesClientToolMapping, error) {
return adaptResponsesClientToolsForFunctionUpstreamWithMapping(
body,
upstream,
apicompat.ResponsesClientToolMapping{},
)
}
func adaptResponsesClientToolsForFunctionUpstreamWithMapping(
body []byte,
upstream string,
inherited apicompat.ResponsesClientToolMapping,
inheritedLoweredTools ...[]any,
) ([]byte, apicompat.ResponsesClientToolMapping, error) {
decoder := json.NewDecoder(bytes.NewReader(body))
decoder.UseNumber()
var requestBody map[string]any
if err := decoder.Decode(&requestBody); err != nil {
return body, apicompat.ResponsesClientToolMapping{}, fmt.Errorf("decode %s Responses client tools: %w", upstream, err)
}
mapping, changed, err := apicompat.AdaptResponsesClientToolsWithInheritedMapping(requestBody, inherited, inheritedLoweredTools...)
if err != nil {
return body, apicompat.ResponsesClientToolMapping{}, err
}
if !changed {
return body, mapping, nil
}
rebuilt, err := marshalOpenAIUpstreamJSON(requestBody)
if err != nil {
return body, apicompat.ResponsesClientToolMapping{}, fmt.Errorf("encode %s Responses client tools: %w", upstream, err)
}
return rebuilt, mapping, nil
}
func adaptGrokResponsesClientTools(body []byte) ([]byte, apicompat.ResponsesClientToolMapping, error) {
return adaptResponsesClientToolsForFunctionUpstream(body, "Grok")
}
func hasResponsesClientToolMapping(mapping apicompat.ResponsesClientToolMapping) bool {
return len(mapping.CustomTools) > 0 || mapping.ToolSearch || len(mapping.NamespaceTools) > 0
}
func hasGrokResponsesClientToolMapping(mapping apicompat.ResponsesClientToolMapping) bool {
return hasResponsesClientToolMapping(mapping)
}
func setGrokResponsesClientToolMapping(c *gin.Context, mapping apicompat.ResponsesClientToolMapping) {
if c == nil {
return
}
if !hasGrokResponsesClientToolMapping(mapping) {
clearGrokResponsesClientToolMapping(c)
return
}
c.Set(grokResponsesClientToolMappingContextKey, mapping)
}
func clearGrokResponsesClientToolMapping(c *gin.Context) {
if c == nil {
return
}
if _, exists := c.Get(grokResponsesClientToolMappingContextKey); !exists {
return
}
c.Set(grokResponsesClientToolMappingContextKey, apicompat.ResponsesClientToolMapping{})
}
func grokResponsesClientToolMapping(c *gin.Context) (apicompat.ResponsesClientToolMapping, bool) {
if c == nil {
return apicompat.ResponsesClientToolMapping{}, false
}
value, ok := c.Get(grokResponsesClientToolMappingContextKey)
if !ok {
return apicompat.ResponsesClientToolMapping{}, false
}
mapping, ok := value.(apicompat.ResponsesClientToolMapping)
return mapping, ok && hasGrokResponsesClientToolMapping(mapping)
}
func restoreGrokResponsesClientToolPayload(c *gin.Context, payload []byte) ([]byte, error) {
mapping, ok := grokResponsesClientToolMapping(c)
if !ok || !bytes.Contains(payload, []byte(`"function_call"`)) || !json.Valid(payload) {
return payload, nil
}
restored, _, err := apicompat.RestoreResponsesClientToolPayload(payload, mapping)
return restored, err
}
type responsesClientToolStreamBody struct {
*io.PipeReader
source io.Closer
}
func (b *responsesClientToolStreamBody) Close() error {
readerErr := b.PipeReader.Close()
sourceErr := b.source.Close()
if readerErr != nil {
return readerErr
}
return sourceErr
}
func newResponsesClientToolStreamBody(
source io.ReadCloser,
mapping apicompat.ResponsesClientToolMapping,
maxLineSize int,
) io.ReadCloser {
reader, writer := io.Pipe()
body := &responsesClientToolStreamBody{PipeReader: reader, source: source}
go transformResponsesClientToolStream(source, writer, mapping, maxLineSize)
return body
}
func newGrokResponsesClientToolStreamBody(
source io.ReadCloser,
mapping apicompat.ResponsesClientToolMapping,
maxLineSize int,
) io.ReadCloser {
return newResponsesClientToolStreamBody(source, mapping, maxLineSize)
}
func transformResponsesClientToolStream(
source io.ReadCloser,
destination *io.PipeWriter,
mapping apicompat.ResponsesClientToolMapping,
maxLineSize int,
) {
defer func() { _ = source.Close() }()
if maxLineSize <= 0 {
maxLineSize = defaultMaxLineSize
}
scanner := bufio.NewScanner(source)
scanBuf := getSSEScannerBuf64K()
defer putSSEScannerBuf64K(scanBuf)
scanner.Buffer(scanBuf[:0], maxLineSize)
documents := newOpenAISSEJSONDocumentScanner(scanner)
restorer := apicompat.NewResponsesClientToolStreamRestorer(mapping)
buffered := bufio.NewWriterSize(destination, 4*1024)
pendingFields := make([]string, 0, 2)
frameHadEventField := false
frameEmitted := false
writeLine := func(line string) error {
if _, err := buffered.WriteString(line); err != nil {
return err
}
return buffered.WriteByte('\n')
}
writePendingFields := func(payload []byte, includeNonEvent bool) error {
eventType := strings.TrimSpace(gjson.GetBytes(payload, "type").String())
for _, field := range pendingFields {
if _, isEvent := extractOpenAISSEEventLine(field); isEvent {
if eventType != "" {
if err := writeLine("event: " + eventType); err != nil {
return err
}
} else if err := writeLine(field); err != nil {
return err
}
continue
}
if includeNonEvent {
if err := writeLine(field); err != nil {
return err
}
}
}
return nil
}
writePayloads := func(payloads [][]byte) error {
for index, payload := range payloads {
if index == 0 {
if err := writePendingFields(payload, true); err != nil {
return err
}
} else if frameHadEventField {
eventType := strings.TrimSpace(gjson.GetBytes(payload, "type").String())
if eventType != "" {
if err := writeLine("event: " + eventType); err != nil {
return err
}
}
}
if err := writeLine("data: " + string(payload)); err != nil {
return err
}
if err := writeLine(""); err != nil {
return err
}
}
return buffered.Flush()
}
for documents.Scan() {
line := documents.Text()
data, isData := extractOpenAISSEDataLine(line)
if isData {
payload := []byte(data)
payloads := [][]byte{payload}
if json.Valid(payload) {
var err error
payloads, _, err = restorer.RestoreEvent(payload)
if err != nil {
_ = buffered.Flush()
_ = destination.CloseWithError(fmt.Errorf("restore Responses client tool event: %w", err))
return
}
}
if err := writePayloads(payloads); err != nil {
_ = destination.CloseWithError(err)
return
}
pendingFields = pendingFields[:0]
frameHadEventField = false
frameEmitted = true
continue
}
if line == "" {
if !frameEmitted {
for _, field := range pendingFields {
if err := writeLine(field); err != nil {
_ = destination.CloseWithError(err)
return
}
}
if len(pendingFields) > 0 {
if err := writeLine(""); err != nil {
_ = destination.CloseWithError(err)
return
}
if err := buffered.Flush(); err != nil {
_ = destination.CloseWithError(err)
return
}
}
}
pendingFields = pendingFields[:0]
frameHadEventField = false
frameEmitted = false
continue
}
if _, isEvent := extractOpenAISSEEventLine(line); isEvent {
frameHadEventField = true
}
pendingFields = append(pendingFields, line)
}
for _, field := range pendingFields {
if err := writeLine(field); err != nil {
_ = destination.CloseWithError(err)
return
}
}
if err := buffered.Flush(); err != nil {
_ = destination.CloseWithError(err)
return
}
if err := documents.Err(); err != nil {
_ = destination.CloseWithError(err)
return
}
_ = destination.Close()
}