Sub2API v1.0 - AI API 网关(二开初始版本,基于上游 Wei-Shaw/sub2api)
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
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
This commit is contained in:
@@ -0,0 +1,284 @@
|
||||
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()
|
||||
}
|
||||
Reference in New Issue
Block a user