Files
sub2api/backend/internal/service/openai_ws_forwarder_logutil.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

692 lines
18 KiB
Go

package service
import (
"bytes"
"context"
"encoding/json"
"errors"
"fmt"
"io"
"net"
"net/http"
"sort"
"strings"
"github.com/Wei-Shaw/sub2api/internal/pkg/logger"
coderws "github.com/coder/websocket"
"github.com/gin-gonic/gin"
"github.com/tidwall/gjson"
"go.uber.org/zap"
)
func normalizeOpenAIWSLogValue(value string) string {
trimmed := strings.TrimSpace(value)
if trimmed == "" {
return "-"
}
return openAIWSLogValueReplacer.Replace(trimmed)
}
func truncateOpenAIWSLogValue(value string, maxLen int) string {
normalized := normalizeOpenAIWSLogValue(value)
if normalized == "-" || maxLen <= 0 {
return normalized
}
if len(normalized) <= maxLen {
return normalized
}
return normalized[:maxLen] + "..."
}
func openAIWSHeaderValueForLog(headers http.Header, key string) string {
if headers == nil {
return "-"
}
return truncateOpenAIWSLogValue(headers.Get(key), openAIWSHeaderValueMaxLen)
}
func hasOpenAIWSHeader(headers http.Header, key string) bool {
if headers == nil {
return false
}
return strings.TrimSpace(headers.Get(key)) != ""
}
type openAIWSSessionHeaderResolution struct {
SessionID string
ConversationID string
SessionSource string
ConversationSource string
}
func resolveOpenAIWSSessionHeaders(c *gin.Context, promptCacheKey string) openAIWSSessionHeaderResolution {
resolution := openAIWSSessionHeaderResolution{
SessionSource: "none",
ConversationSource: "none",
}
if c != nil && c.Request != nil {
if sessionID := strings.TrimSpace(c.Request.Header.Get("session_id")); sessionID != "" {
resolution.SessionID = sessionID
resolution.SessionSource = "header_session_id"
}
if conversationID := strings.TrimSpace(c.Request.Header.Get("conversation_id")); conversationID != "" {
resolution.ConversationID = conversationID
resolution.ConversationSource = "header_conversation_id"
if resolution.SessionID == "" {
resolution.SessionID = conversationID
resolution.SessionSource = "header_conversation_id"
}
}
}
cacheKey := strings.TrimSpace(promptCacheKey)
if cacheKey != "" {
if resolution.SessionID == "" {
resolution.SessionID = cacheKey
resolution.SessionSource = "prompt_cache_key"
}
}
return resolution
}
func shouldLogOpenAIWSEvent(idx int, eventType string) bool {
if idx <= openAIWSEventLogHeadLimit {
return true
}
if openAIWSEventLogEveryN > 0 && idx%openAIWSEventLogEveryN == 0 {
return true
}
if eventType == "error" || isOpenAIWSTerminalEvent(eventType) {
return true
}
return false
}
func shouldLogOpenAIWSBufferedEvent(idx int) bool {
if idx <= openAIWSBufferLogHeadLimit {
return true
}
if openAIWSBufferLogEveryN > 0 && idx%openAIWSBufferLogEveryN == 0 {
return true
}
return false
}
func openAIWSEventMayContainModel(eventType string) bool {
switch eventType {
case "response.created",
"response.in_progress",
"response.completed",
"response.done",
"response.failed",
"response.incomplete",
"response.cancelled",
"response.canceled":
return true
default:
trimmed := strings.TrimSpace(eventType)
if trimmed == eventType {
return false
}
switch trimmed {
case "response.created",
"response.in_progress",
"response.completed",
"response.done",
"response.failed",
"response.incomplete",
"response.cancelled",
"response.canceled":
return true
default:
return false
}
}
}
func openAIWSEventMayContainToolCalls(eventType string) bool {
eventType = strings.TrimSpace(eventType)
if eventType == "" {
return false
}
if strings.Contains(eventType, "function_call") || strings.Contains(eventType, "tool_call") {
return true
}
switch eventType {
case "response.output_item.added", "response.output_item.done", "response.completed", "response.done":
return true
default:
return false
}
}
func openAIWSEventShouldParseUsage(eventType string) bool {
switch strings.TrimSpace(eventType) {
case "response.completed", "response.done", "response.failed", "response.incomplete", "response.cancelled", "response.canceled":
return true
default:
return false
}
}
func parseOpenAIWSEventEnvelope(message []byte) (eventType string, responseID string, response gjson.Result) {
if len(message) == 0 {
return "", "", gjson.Result{}
}
values := gjson.GetManyBytes(message, "type", "response.id", "id", "response")
eventType = strings.TrimSpace(values[0].String())
if id := strings.TrimSpace(values[1].String()); id != "" {
responseID = id
} else {
responseID = strings.TrimSpace(values[2].String())
}
return eventType, responseID, values[3]
}
func openAIWSMessageLikelyContainsToolCalls(message []byte) bool {
if len(message) == 0 {
return false
}
return bytes.Contains(message, []byte(`"tool_calls"`)) ||
bytes.Contains(message, []byte(`"tool_call"`)) ||
bytes.Contains(message, []byte(`"function_call"`))
}
func parseOpenAIWSResponseUsageFromCompletedEvent(message []byte, usage *OpenAIUsage) {
if usage == nil || len(message) == 0 {
return
}
if parsedUsage, ok := extractOpenAIUsageFromJSONBytes(message); ok {
*usage = parsedUsage
}
}
func parseOpenAIWSErrorEventFields(message []byte) (code string, errType string, errMessage string) {
if len(message) == 0 {
return "", "", ""
}
values := gjson.GetManyBytes(message, "error.code", "error.type", "error.message")
return strings.TrimSpace(values[0].String()), strings.TrimSpace(values[1].String()), strings.TrimSpace(values[2].String())
}
func summarizeOpenAIWSErrorEventFieldsFromRaw(codeRaw, errTypeRaw, errMessageRaw string) (code string, errType string, errMessage string) {
code = truncateOpenAIWSLogValue(codeRaw, openAIWSLogValueMaxLen)
errType = truncateOpenAIWSLogValue(errTypeRaw, openAIWSLogValueMaxLen)
errMessage = truncateOpenAIWSLogValue(errMessageRaw, openAIWSLogValueMaxLen)
return code, errType, errMessage
}
func summarizeOpenAIWSErrorEventFields(message []byte) (code string, errType string, errMessage string) {
if len(message) == 0 {
return "-", "-", "-"
}
return summarizeOpenAIWSErrorEventFieldsFromRaw(parseOpenAIWSErrorEventFields(message))
}
func summarizeOpenAIWSPayloadKeySizes(payload map[string]any, topN int) string {
if len(payload) == 0 {
return "-"
}
type keySize struct {
Key string
Size int
}
sizes := make([]keySize, 0, len(payload))
for key, value := range payload {
size := estimateOpenAIWSPayloadValueSize(value, openAIWSPayloadSizeEstimateDepth)
sizes = append(sizes, keySize{Key: key, Size: size})
}
sort.Slice(sizes, func(i, j int) bool {
if sizes[i].Size == sizes[j].Size {
return sizes[i].Key < sizes[j].Key
}
return sizes[i].Size > sizes[j].Size
})
if topN <= 0 || topN > len(sizes) {
topN = len(sizes)
}
parts := make([]string, 0, topN)
for idx := 0; idx < topN; idx++ {
item := sizes[idx]
parts = append(parts, fmt.Sprintf("%s:%d", item.Key, item.Size))
}
return strings.Join(parts, ",")
}
func estimateOpenAIWSPayloadValueSize(value any, depth int) int {
if depth <= 0 {
return -1
}
switch v := value.(type) {
case nil:
return 0
case string:
return len(v)
case []byte:
return len(v)
case bool:
return 1
case int, int8, int16, int32, int64, uint, uint8, uint16, uint32, uint64:
return 8
case float32, float64:
return 8
case map[string]any:
if len(v) == 0 {
return 2
}
total := 2
count := 0
for key, item := range v {
count++
if count > openAIWSPayloadSizeEstimateMaxItems {
return -1
}
itemSize := estimateOpenAIWSPayloadValueSize(item, depth-1)
if itemSize < 0 {
return -1
}
total += len(key) + itemSize + 3
if total > openAIWSPayloadSizeEstimateMaxBytes {
return -1
}
}
return total
case []any:
if len(v) == 0 {
return 2
}
total := 2
limit := len(v)
if limit > openAIWSPayloadSizeEstimateMaxItems {
return -1
}
for i := 0; i < limit; i++ {
itemSize := estimateOpenAIWSPayloadValueSize(v[i], depth-1)
if itemSize < 0 {
return -1
}
total += itemSize + 1
if total > openAIWSPayloadSizeEstimateMaxBytes {
return -1
}
}
return total
default:
raw, err := json.Marshal(v)
if err != nil {
return -1
}
if len(raw) > openAIWSPayloadSizeEstimateMaxBytes {
return -1
}
return len(raw)
}
}
func openAIWSPayloadString(payload map[string]any, key string) string {
if len(payload) == 0 {
return ""
}
raw, ok := payload[key]
if !ok {
return ""
}
switch v := raw.(type) {
case nil:
return ""
case string:
return strings.TrimSpace(v)
case []byte:
return strings.TrimSpace(string(v))
default:
return ""
}
}
func openAIWSPayloadStringFromRaw(payload []byte, key string) string {
if len(payload) == 0 || strings.TrimSpace(key) == "" {
return ""
}
return strings.TrimSpace(gjson.GetBytes(payload, key).String())
}
func openAIWSPayloadBoolFromRaw(payload []byte, key string, defaultValue bool) bool {
if len(payload) == 0 || strings.TrimSpace(key) == "" {
return defaultValue
}
value := gjson.GetBytes(payload, key)
if !value.Exists() {
return defaultValue
}
if value.Type != gjson.True && value.Type != gjson.False {
return defaultValue
}
return value.Bool()
}
func openAIWSSessionHashesFromID(sessionID string) (string, string) {
return deriveOpenAISessionHashes(sessionID)
}
func extractOpenAIWSImageURL(value any) string {
switch v := value.(type) {
case string:
return strings.TrimSpace(v)
case map[string]any:
if raw, ok := v["url"].(string); ok {
return strings.TrimSpace(raw)
}
}
return ""
}
func summarizeOpenAIWSInput(input any) string {
items, ok := input.([]any)
if !ok || len(items) == 0 {
return "-"
}
itemCount := len(items)
textChars := 0
imageDataURLs := 0
imageDataURLChars := 0
imageRemoteURLs := 0
handleContentItem := func(contentItem map[string]any) {
contentType, _ := contentItem["type"].(string)
switch strings.TrimSpace(contentType) {
case "input_text", "output_text", "text":
if text, ok := contentItem["text"].(string); ok {
textChars += len(text)
}
case "input_image":
imageURL := extractOpenAIWSImageURL(contentItem["image_url"])
if imageURL == "" {
return
}
if strings.HasPrefix(strings.ToLower(imageURL), "data:image/") {
imageDataURLs++
imageDataURLChars += len(imageURL)
return
}
imageRemoteURLs++
}
}
handleInputItem := func(inputItem map[string]any) {
if content, ok := inputItem["content"].([]any); ok {
for _, rawContent := range content {
contentItem, ok := rawContent.(map[string]any)
if !ok {
continue
}
handleContentItem(contentItem)
}
return
}
itemType, _ := inputItem["type"].(string)
switch strings.TrimSpace(itemType) {
case "input_text", "output_text", "text":
if text, ok := inputItem["text"].(string); ok {
textChars += len(text)
}
case "input_image":
imageURL := extractOpenAIWSImageURL(inputItem["image_url"])
if imageURL == "" {
return
}
if strings.HasPrefix(strings.ToLower(imageURL), "data:image/") {
imageDataURLs++
imageDataURLChars += len(imageURL)
return
}
imageRemoteURLs++
}
}
for _, rawItem := range items {
inputItem, ok := rawItem.(map[string]any)
if !ok {
continue
}
handleInputItem(inputItem)
}
return fmt.Sprintf(
"items=%d,text_chars=%d,image_data_urls=%d,image_data_url_chars=%d,image_remote_urls=%d",
itemCount,
textChars,
imageDataURLs,
imageDataURLChars,
imageRemoteURLs,
)
}
func dropOpenAIWSPayloadKey(payload map[string]any, key string, removed *[]string) {
if len(payload) == 0 || strings.TrimSpace(key) == "" {
return
}
if _, exists := payload[key]; !exists {
return
}
delete(payload, key)
*removed = append(*removed, key)
}
// applyOpenAIWSRetryPayloadStrategy 在 WS 连续失败时仅移除无语义字段,
// 避免重试成功却改变原始请求语义。
// 注意:prompt_cache_key 不应在重试中移除;它常用于会话稳定标识(session_id 兜底)。
func applyOpenAIWSRetryPayloadStrategy(payload map[string]any, attempt int) (strategy string, removedKeys []string) {
if len(payload) == 0 {
return "empty", nil
}
if attempt <= 1 {
return "full", nil
}
removed := make([]string, 0, 2)
if attempt >= 2 {
dropOpenAIWSPayloadKey(payload, "include", &removed)
}
if len(removed) == 0 {
return "full", nil
}
sort.Strings(removed)
return "trim_optional_fields", removed
}
func logOpenAIWSModeInfo(format string, args ...any) {
logger.LegacyPrintf("service.openai_gateway", "[OpenAI WS Mode][openai_ws_mode=true] "+format, args...)
}
func isOpenAIWSModeDebugEnabled() bool {
return logger.L().Core().Enabled(zap.DebugLevel)
}
func logOpenAIWSModeDebug(format string, args ...any) {
if !isOpenAIWSModeDebugEnabled() {
return
}
logger.LegacyPrintf("service.openai_gateway", "[debug] [OpenAI WS Mode][openai_ws_mode=true] "+format, args...)
}
func logOpenAIWSBindResponseAccountWarn(groupID, accountID int64, responseID string, err error) {
if err == nil {
return
}
logger.L().Warn(
"openai.ws_bind_response_account_failed",
zap.Int64("group_id", groupID),
zap.Int64("account_id", accountID),
zap.String("response_id", truncateOpenAIWSLogValue(responseID, openAIWSIDValueMaxLen)),
zap.Error(err),
)
}
func summarizeOpenAIWSReadCloseError(err error) (status string, reason string) {
if err == nil {
return "-", "-"
}
statusCode := coderws.CloseStatus(err)
if statusCode == -1 {
return "-", "-"
}
closeStatus := fmt.Sprintf("%d(%s)", int(statusCode), statusCode.String())
closeReason := "-"
var closeErr coderws.CloseError
if errors.As(err, &closeErr) {
reasonText := strings.TrimSpace(closeErr.Reason)
if reasonText != "" {
closeReason = normalizeOpenAIWSLogValue(reasonText)
}
}
return normalizeOpenAIWSLogValue(closeStatus), closeReason
}
func unwrapOpenAIWSDialBaseError(err error) error {
if err == nil {
return nil
}
var dialErr *openAIWSDialError
if errors.As(err, &dialErr) && dialErr != nil && dialErr.Err != nil {
return dialErr.Err
}
return err
}
func openAIWSDialRespHeaderForLog(err error, key string) string {
var dialErr *openAIWSDialError
if !errors.As(err, &dialErr) || dialErr == nil || dialErr.ResponseHeaders == nil {
return "-"
}
return truncateOpenAIWSLogValue(dialErr.ResponseHeaders.Get(key), openAIWSHeaderValueMaxLen)
}
func classifyOpenAIWSDialError(err error) string {
if err == nil {
return "-"
}
baseErr := unwrapOpenAIWSDialBaseError(err)
if baseErr == nil {
return "-"
}
if errors.Is(baseErr, context.DeadlineExceeded) {
return "ctx_deadline_exceeded"
}
if errors.Is(baseErr, context.Canceled) {
return "ctx_canceled"
}
var netErr net.Error
if errors.As(baseErr, &netErr) && netErr.Timeout() {
return "net_timeout"
}
if status := coderws.CloseStatus(baseErr); status != -1 {
return normalizeOpenAIWSLogValue(fmt.Sprintf("ws_close_%d", int(status)))
}
message := strings.ToLower(strings.TrimSpace(baseErr.Error()))
switch {
case strings.Contains(message, "handshake not finished"):
return "handshake_not_finished"
case strings.Contains(message, "bad handshake"):
return "bad_handshake"
case strings.Contains(message, "connection refused"):
return "connection_refused"
case strings.Contains(message, "no such host"):
return "dns_not_found"
case strings.Contains(message, "tls"):
return "tls_error"
case strings.Contains(message, "i/o timeout"):
return "io_timeout"
case strings.Contains(message, "context deadline exceeded"):
return "ctx_deadline_exceeded"
default:
return "dial_error"
}
}
func summarizeOpenAIWSDialError(err error) (
statusCode int,
dialClass string,
closeStatus string,
closeReason string,
respServer string,
respVia string,
respCFRay string,
respRequestID string,
) {
dialClass = "-"
closeStatus = "-"
closeReason = "-"
respServer = "-"
respVia = "-"
respCFRay = "-"
respRequestID = "-"
if err == nil {
return
}
var dialErr *openAIWSDialError
if errors.As(err, &dialErr) && dialErr != nil {
statusCode = dialErr.StatusCode
respServer = openAIWSDialRespHeaderForLog(err, "server")
respVia = openAIWSDialRespHeaderForLog(err, "via")
respCFRay = openAIWSDialRespHeaderForLog(err, "cf-ray")
respRequestID = openAIWSDialRespHeaderForLog(err, "x-request-id")
}
dialClass = normalizeOpenAIWSLogValue(classifyOpenAIWSDialError(err))
closeStatus, closeReason = summarizeOpenAIWSReadCloseError(unwrapOpenAIWSDialBaseError(err))
return
}
func isOpenAIWSClientDisconnectError(err error) bool {
if err == nil {
return false
}
if errors.Is(err, io.EOF) || errors.Is(err, net.ErrClosed) || errors.Is(err, context.Canceled) {
return true
}
switch coderws.CloseStatus(err) {
case coderws.StatusNormalClosure, coderws.StatusGoingAway, coderws.StatusNoStatusRcvd, coderws.StatusAbnormalClosure:
return true
}
message := strings.ToLower(strings.TrimSpace(err.Error()))
if message == "" {
return false
}
return strings.Contains(message, "failed to read frame header: eof") ||
strings.Contains(message, "unexpected eof") ||
strings.Contains(message, "use of closed network connection") ||
strings.Contains(message, "connection reset by peer") ||
strings.Contains(message, "broken pipe") ||
strings.Contains(message, "an existing connection was forcibly closed by the remote host") ||
strings.Contains(message, "an established connection was aborted")
}
func classifyOpenAIWSReadFallbackReason(err error) string {
if err == nil {
return "read_event"
}
switch coderws.CloseStatus(err) {
case coderws.StatusPolicyViolation:
return "policy_violation"
case coderws.StatusMessageTooBig:
return "message_too_big"
default:
return "read_event"
}
}
func sortedKeys(m map[string]any) []string {
if len(m) == 0 {
return nil
}
keys := make([]string, 0, len(m))
for k := range m {
keys = append(keys, k)
}
sort.Strings(keys)
return keys
}