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

855 lines
26 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
package service
import (
"bytes"
"context"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"io"
"net/http"
"net/url"
"path"
"strings"
"time"
"github.com/Wei-Shaw/sub2api/internal/pkg/logger"
coderws "github.com/coder/websocket"
"github.com/google/uuid"
"github.com/tidwall/gjson"
"go.uber.org/zap"
)
const (
defaultLiveMaxSessionDuration = time.Hour
liveLeaseRefreshInterval = 20 * time.Second
liveRedisOperationTimeout = 3 * time.Second
liveClosedRecordTTL = 24 * time.Hour
liveObserverPollInterval = 250 * time.Millisecond
liveObserverStoreRetryLimit = 5
liveUpstreamBodyLimit = 2 << 20
)
// liveObserverStoreRetryInterval 是 var 以便测试缩短 store 报错的重试等待。
var liveObserverStoreRetryInterval = time.Second
var (
chatGPTLiveCallsURL = "https://chatgpt.com/backend-api/codex/realtime/calls?intent=quicksilver&architecture=avas"
chatGPTLiveSidebandBaseURL = "wss://chatgpt.com/backend-api/codex"
)
type liveFrameConn interface {
ReadFrame(ctx context.Context) (coderws.MessageType, []byte, error)
WriteFrame(ctx context.Context, msgType coderws.MessageType, payload []byte) error
Close() error
}
func liveSidebandReadError(err error) error {
if coderws.CloseStatus(err) == coderws.StatusNormalClosure {
return ErrLiveCallNotFound
}
return err
}
func hashLiveCallID(callID string) string {
sum := sha256.Sum256([]byte(callID))
return hex.EncodeToString(sum[:])
}
func liveGroupID(groupID *int64) int64 {
if groupID == nil {
return 0
}
return *groupID
}
func liveOptionalID(value int64) *int64 {
if value <= 0 {
return nil
}
result := value
return &result
}
func (s *OpenAIGatewayService) liveStore() (LiveCallStore, error) {
if s == nil || s.cache == nil {
return nil, ErrLiveUnavailable
}
store, ok := s.cache.(LiveCallStore)
if !ok {
return nil, ErrLiveUnavailable
}
return store, nil
}
func (s *OpenAIGatewayService) liveConcurrencyCache() (LiveConcurrencyCache, error) {
if s == nil || s.concurrencyService == nil || s.concurrencyService.cache == nil {
return nil, ErrLiveUnavailable
}
cache, ok := s.concurrencyService.cache.(LiveConcurrencyCache)
if !ok {
return nil, ErrLiveUnavailable
}
return cache, nil
}
func (s *OpenAIGatewayService) liveMaxSessionDuration() time.Duration {
if s != nil && s.cfg != nil && s.cfg.Gateway.Live.MaxSessionDurationSeconds > 0 {
return time.Duration(s.cfg.Gateway.Live.MaxSessionDurationSeconds) * time.Second
}
return defaultLiveMaxSessionDuration
}
func ValidateLiveCallRequest(request *LiveCallRequest) error {
if request == nil || strings.TrimSpace(request.SDP) == "" {
return errors.New("sdp is required")
}
if len(request.Session) == 0 || !json.Valid(request.Session) {
return errors.New("session must be valid JSON")
}
var sessionObject map[string]json.RawMessage
if err := json.Unmarshal(request.Session, &sessionObject); err != nil {
return errors.New("session must be a JSON object")
}
if sessionObject == nil {
return errors.New("session must be a JSON object")
}
return nil
}
// CreateLiveCall 创建 Frameless 会话。调用方须在调用期间持有普通用户槽位;
// 调度器持有的普通账号槽位会被同一个 Live 租约原子接替。
func (s *OpenAIGatewayService) CreateLiveCall(
ctx context.Context,
request *LiveCallRequest,
identity LiveCallIdentity,
userMaxConcurrency int,
) (*LiveCallCreated, error) {
if err := ValidateLiveCallRequest(request); err != nil {
return nil, err
}
store, err := s.liveStore()
if err != nil {
return nil, err
}
liveCache, err := s.liveConcurrencyCache()
if err != nil {
return nil, err
}
attestation, attestationCiphertext, err := s.prepareLiveAttestation(ctx)
if err != nil {
return nil, err
}
excluded := make(map[int64]struct{})
// Live 按通话时长计费,不属于 token 利润门的语义范围:显式豁免,避免
// 防御性装门按文本 D 过滤 Live 账号池且门与计费时刻不同源。
ctx = WithOpenAIProfitControlSuppressed(ctx)
var lastErr error
for attempt := 0; attempt <= 3; attempt++ {
selection, _, selectErr := s.SelectAccountWithSchedulerForCapability(
ctx,
identity.GroupID,
"",
uuid.NewString(),
"",
excluded,
OpenAIUpstreamTransportHTTPSSE,
OpenAIEndpointCapabilityLive,
false,
false,
false,
)
if selectErr != nil {
if lastErr != nil {
return nil, lastErr
}
return nil, selectErr
}
if selection == nil || selection.Account == nil || !selection.Acquired {
if selection != nil && selection.ReleaseFunc != nil {
selection.ReleaseFunc()
}
return nil, ErrLiveConcurrencyFull
}
account := selection.Account
leaseID := generateRequestID()
acquired, acquireErr := liveCache.AcquireLiveLease(
ctx,
account.ID,
account.Concurrency,
identity.UserID,
userMaxConcurrency,
identity.APIKeyID,
leaseID,
true,
)
if acquireErr != nil || !acquired {
selection.ReleaseFunc()
if acquireErr != nil {
return nil, acquireErr
}
return nil, ErrLiveConcurrencyFull
}
created, createErr := s.createUpstreamLiveCall(ctx, account, request, attestation)
selection.ReleaseFunc()
if createErr != nil {
s.releaseLiveLease(account.ID, identity.UserID, identity.APIKeyID, leaseID)
if !s.shouldFailoverLiveCreateError(createErr) {
return nil, createErr
}
excluded[account.ID] = struct{}{}
lastErr = createErr
continue
}
now := time.Now()
model := strings.TrimSpace(gjson.GetBytes(request.Session, "model").String())
if model == "" {
model = "gpt-live"
}
record := &LiveCallRecord{
CallID: created.CallID,
CallHash: hashLiveCallID(created.CallID),
AccountID: account.ID,
APIKeyID: identity.APIKeyID,
UserID: identity.UserID,
GroupID: liveGroupID(identity.GroupID),
SubscriptionID: liveGroupID(identity.SubscriptionID),
LeaseID: leaseID,
Model: model,
CreatedAt: now,
ExpiresAt: now.Add(s.liveMaxSessionDuration()),
Controller: LiveControllerPending,
UserAgent: identity.UserAgent,
IPAddress: identity.IPAddress,
InboundEndpoint: identity.InboundEndpoint,
AttestationCiphertext: attestationCiphertext,
}
mappingTTL := s.liveMaxSessionDuration() + 5*time.Minute
if saveErr := store.SaveLiveCall(ctx, record, mappingTTL); saveErr != nil {
s.releaseLiveLease(account.ID, identity.UserID, identity.APIKeyID, leaseID)
return nil, fmt.Errorf("save live call mapping: %w", saveErr)
}
created.Account = account
go s.observeLiveCall(record)
return created, nil
}
if lastErr != nil {
return nil, lastErr
}
return nil, ErrLiveUnavailable
}
func (s *OpenAIGatewayService) shouldFailoverLiveCreateError(err error) bool {
var upstreamErr *UpstreamFailoverError
if !errors.As(err, &upstreamErr) {
// 凭证读取和网络传输错误都可能只影响当前账号或代理。
return true
}
return s.shouldFailoverOpenAIUpstreamResponse(
upstreamErr.StatusCode,
"",
upstreamErr.ResponseBody,
)
}
func (s *OpenAIGatewayService) createUpstreamLiveCall(
ctx context.Context,
account *Account,
request *LiveCallRequest,
attestation string,
) (*LiveCallCreated, error) {
token, _, err := s.GetAccessToken(ctx, account)
if err != nil {
logLiveCreateStageFailure(ctx, account.ID, "access_token", err)
return nil, err
}
body, err := json.Marshal(struct {
SDP string `json:"sdp"`
Session json.RawMessage `json:"session"`
}{
SDP: request.SDP,
Session: request.Session,
})
if err != nil {
return nil, err
}
reqCtx := WithHTTPUpstreamRedirectsDisabled(WithHTTPUpstreamProfile(ctx, HTTPUpstreamProfileOpenAI))
upstreamReq, err := http.NewRequestWithContext(reqCtx, http.MethodPost, chatGPTLiveCallsURL, bytes.NewReader(body))
if err != nil {
return nil, err
}
authHeaders, err := s.buildOpenAIAuthenticationHeaders(ctx, account, token)
if err != nil {
logLiveCreateStageFailure(ctx, account.ID, "authentication_headers", err)
return nil, err
}
for key, values := range authHeaders {
for _, value := range values {
upstreamReq.Header.Add(key, value)
}
}
upstreamReq.Host = "chatgpt.com"
if err := resolveAndSetOpenAIChatGPTAccountHeaders(ctx, s.accountRepo, upstreamReq.Header, account); err != nil {
logLiveCreateStageFailure(ctx, account.ID, "account_headers", err)
return nil, err
}
upstreamReq.Header.Set("Content-Type", "application/json")
upstreamReq.Header.Set("Accept", "application/sdp")
upstreamReq.Header.Set(liveAttestationHeader, attestation)
applyLiveUpstreamIdentityHeaders(upstreamReq.Header)
resp, err := s.httpUpstream.Do(upstreamReq, resolveAccountProxyURL(account), account.ID, account.Concurrency)
if err != nil {
logLiveCreateStageFailure(ctx, account.ID, "upstream_transport", err)
return nil, err
}
defer func() { _ = resp.Body.Close() }()
responseBody, readErr := io.ReadAll(io.LimitReader(resp.Body, liveUpstreamBodyLimit+1))
if readErr != nil {
return nil, readErr
}
if len(responseBody) > liveUpstreamBodyLimit {
return nil, errors.New("live upstream response is too large")
}
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
logLiveUpstreamFailure(ctx, account.ID, resp.StatusCode, resp.Header, responseBody)
return nil, &UpstreamFailoverError{
StatusCode: resp.StatusCode,
ResponseBody: responseBody,
ResponseHeaders: resp.Header.Clone(),
}
}
callID, err := liveCallIDFromLocation(resp.Header.Get("Location"))
if err != nil {
return nil, err
}
return &LiveCallCreated{
SDP: responseBody,
CallID: callID,
Location: resp.Header.Get("Location"),
}, nil
}
func logLiveCreateStageFailure(ctx context.Context, accountID int64, stage string, err error) {
logger.FromContext(ctx).Warn(
"OpenAI Live 创建阶段失败",
zap.Int64("account_id", accountID),
zap.String("stage", stage),
zap.String("error_type", fmt.Sprintf("%T", err)),
)
}
func logLiveUpstreamFailure(
ctx context.Context,
accountID int64,
statusCode int,
headers http.Header,
body []byte,
) {
errorType := strings.TrimSpace(gjson.GetBytes(body, "error.type").String())
errorCode := strings.TrimSpace(gjson.GetBytes(body, "error.code").String())
errorMessage := strings.TrimSpace(gjson.GetBytes(body, "error.message").String())
if errorType == "" {
errorType = strings.TrimSpace(gjson.GetBytes(body, "type").String())
}
if errorCode == "" {
errorCode = strings.TrimSpace(gjson.GetBytes(body, "code").String())
}
if errorMessage == "" {
errorMessage = strings.TrimSpace(gjson.GetBytes(body, "message").String())
}
if errorMessage == "" {
errorMessage = strings.TrimSpace(gjson.GetBytes(body, "detail").String())
}
logger.FromContext(ctx).Warn(
"OpenAI Live 上游拒绝请求",
zap.Int64("account_id", accountID),
zap.Int("upstream_status_code", statusCode),
zap.String("upstream_error_type", truncateOpenAIWSLogValue(errorType, 120)),
zap.String("upstream_error_code", truncateOpenAIWSLogValue(errorCode, 120)),
zap.String("upstream_error_message", truncateOpenAIWSLogValue(errorMessage, 300)),
zap.String("upstream_content_type", truncateOpenAIWSLogValue(headers.Get("Content-Type"), 120)),
zap.String("upstream_server", truncateOpenAIWSLogValue(headers.Get("Server"), 120)),
zap.String("upstream_cf_mitigated", truncateOpenAIWSLogValue(headers.Get("Cf-Mitigated"), 120)),
zap.String("upstream_cf_ray", truncateOpenAIWSLogValue(headers.Get("Cf-Ray"), 120)),
zap.String("upstream_request_id", truncateOpenAIWSLogValue(headers.Get("X-Request-Id"), 120)),
)
}
func liveCallIDFromLocation(location string) (string, error) {
location = strings.TrimSpace(location)
if location == "" {
return "", errors.New("live upstream response has no Location")
}
parsed, err := url.Parse(location)
if err != nil {
return "", fmt.Errorf("parse live Location: %w", err)
}
callID := strings.TrimSpace(path.Base(strings.TrimSuffix(parsed.Path, "/")))
if callID == "" || callID == "." || callID == "codex" {
return "", errors.New("live upstream Location has no call id")
}
return callID, nil
}
func applyLiveUpstreamIdentityHeaders(headers http.Header) {
headers.Set("OpenAI-Alpha", "quicksilver=v2")
ensureCodexIdentityHeaders(headers)
enforceCodexIdentityHeaders(headers)
if strings.TrimSpace(headers.Get("session-id")) == "" {
headers.Set("session-id", uuid.NewString())
}
if strings.TrimSpace(headers.Get("thread-id")) == "" {
headers.Set("thread-id", uuid.NewString())
}
// Realtime/Live 不使用 Responses 的实验头。
headers.Del("OpenAI-Beta")
}
func (s *OpenAIGatewayService) liveSidebandHeaders(
ctx context.Context,
account *Account,
record *LiveCallRecord,
) (http.Header, error) {
token, _, err := s.GetAccessToken(ctx, account)
if err != nil {
return nil, err
}
headers, err := s.buildOpenAIAuthenticationHeaders(ctx, account, token)
if err != nil {
return nil, err
}
if err := resolveAndSetOpenAIChatGPTAccountHeaders(ctx, s.accountRepo, headers, account); err != nil {
return nil, err
}
attestation, err := s.decryptLiveAttestation(record)
if err != nil {
return nil, err
}
headers.Set(liveAttestationHeader, attestation)
applyLiveUpstreamIdentityHeaders(headers)
return headers, nil
}
func (s *OpenAIGatewayService) dialLiveSideband(ctx context.Context, record *LiveCallRecord) (liveFrameConn, error) {
account, err := s.accountRepo.GetByID(ctx, record.AccountID)
if err != nil {
return nil, err
}
if account == nil || !account.SupportsOpenAIEndpointCapability(OpenAIEndpointCapabilityLive) {
return nil, ErrLiveUnavailable
}
headers, err := s.liveSidebandHeaders(ctx, account, record)
if err != nil {
return nil, err
}
target := strings.TrimRight(chatGPTLiveSidebandBaseURL, "/") + "/" + url.PathEscape(record.CallID)
conn, status, _, err := s.getOpenAIWSPassthroughDialer().Dial(ctx, target, headers, resolveAccountProxyURL(account))
if err != nil {
return nil, fmt.Errorf("dial live sideband (status %d): %w", status, err)
}
raw, ok := conn.(liveFrameConn)
if !ok {
_ = conn.Close()
return nil, errors.New("live sideband transport does not support raw frames")
}
return raw, nil
}
func (s *OpenAIGatewayService) GetLiveCallForIdentity(
ctx context.Context,
callID string,
identity LiveCallIdentity,
) (*LiveCallRecord, error) {
store, err := s.liveStore()
if err != nil {
return nil, err
}
record, err := store.GetLiveCall(ctx, hashLiveCallID(callID))
if err != nil {
return nil, err
}
if record.CallID != callID ||
record.APIKeyID != identity.APIKeyID ||
record.UserID != identity.UserID ||
record.GroupID != liveGroupID(identity.GroupID) {
return nil, ErrLiveIdentityMismatch
}
if record.Controller == LiveControllerClosed {
return nil, ErrLiveCallNotFound
}
return record, nil
}
// ProxyLiveSideband 让认证后的客户端接管控制连接;媒体始终不经过这里。
func (s *OpenAIGatewayService) ProxyLiveSideband(
ctx context.Context,
record *LiveCallRecord,
downstream *coderws.Conn,
) error {
if record == nil || downstream == nil {
return ErrLiveCallNotFound
}
store, err := s.liveStore()
if err != nil {
return err
}
owner := uuid.NewString()
claimed, err := store.ClaimLiveController(ctx, record.CallHash, LiveControllerProxy, owner)
if err != nil {
return err
}
if !claimed {
return ErrLiveControllerChanged
}
// observer 轮询到接管状态后会关闭旧控制连接;同一个 call 可重新加入。
time.Sleep(liveObserverPollInterval)
upstream, err := s.dialLiveSideband(ctx, record)
if err != nil {
_, _ = store.ReleaseLiveController(context.Background(), record.CallHash, owner)
go s.observeLiveCall(record)
return err
}
defer func() { _ = upstream.Close() }()
downstream.SetReadLimit(openAIWSMessageReadLimitBytes)
proxyCtx, cancel := context.WithCancel(ctx)
defer cancel()
errCh := make(chan error, 2)
go func() {
for {
messageType, payload, readErr := downstream.Read(proxyCtx)
if readErr != nil {
errCh <- readErr
return
}
if writeErr := upstream.WriteFrame(proxyCtx, messageType, payload); writeErr != nil {
errCh <- writeErr
return
}
}
}()
go func() {
for {
messageType, payload, readErr := upstream.ReadFrame(proxyCtx)
if readErr != nil {
errCh <- liveSidebandReadError(readErr)
return
}
if writeErr := downstream.Write(proxyCtx, messageType, payload); writeErr != nil {
errCh <- writeErr
return
}
if messageType == coderws.MessageText {
eventType := strings.TrimSpace(gjson.GetBytes(payload, "type").String())
if eventType == "session.closed" || eventType == "session.ended" {
errCh <- ErrLiveCallNotFound
return
}
}
}
}()
runErr := s.runLiveController(proxyCtx, record, upstream, errCh)
cancel()
_, _ = store.ReleaseLiveController(context.Background(), record.CallHash, owner)
if liveSessionEnded(runErr) || !time.Now().Before(record.ExpiresAt) {
s.finalizeLiveCall(record)
return runErr
}
go s.observeLiveCall(record)
return runErr
}
// liveSessionEnded 判断控制连接的退出原因是否意味着会话已终结(应 finalize:写
// usage log 并释放租约),而不是可以交给 observer 重连的临时错误。
//
// ErrLiveUnavailable 在控制循环里只会来自租约续租失败。RefreshLiveLease 的 Lua 在
// leaseID 被 GC 后不会重新写入,重连也拿不回并发槽 —— 若按临时错误重试,会话会以
// 约 1 秒一轮的节奏空转到 ExpiresAt,期间持着上游连接却不计入任何并发限制。
func liveSessionEnded(err error) bool {
return errors.Is(err, ErrLiveCallNotFound) ||
errors.Is(err, ErrLiveUnavailable) ||
errors.Is(err, context.DeadlineExceeded)
}
func (s *OpenAIGatewayService) runLiveController(
ctx context.Context,
record *LiveCallRecord,
upstream liveFrameConn,
errCh <-chan error,
) error {
refreshTicker := time.NewTicker(liveLeaseRefreshInterval)
defer refreshTicker.Stop()
maxTimer := time.NewTimer(time.Until(record.ExpiresAt))
defer maxTimer.Stop()
for {
select {
case <-ctx.Done():
return context.Cause(ctx)
case err := <-errCh:
return err
case <-maxTimer.C:
closeCtx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
_ = upstream.WriteFrame(closeCtx, coderws.MessageText, []byte(`{"type":"session.close"}`))
cancel()
return context.DeadlineExceeded
case <-refreshTicker.C:
if !s.refreshLiveLease(record) {
return ErrLiveUnavailable
}
}
}
}
func (s *OpenAIGatewayService) observeLiveCall(record *LiveCallRecord) {
if record == nil {
return
}
store, err := s.liveStore()
if err != nil {
return
}
owner := uuid.NewString()
claimed, claimErr := store.ClaimLiveController(context.Background(), record.CallHash, LiveControllerObserver, owner)
if claimErr != nil {
// store 报错时无法确认控制权归属,不能静默退出:若 claim 实际已生效而
// observer 消失,租约与 usage log 都会丢。兜底 finalize 是幂等的,即使
// 控制权在他人手上也只会在到期后落一次库。
s.finalizeLiveCallAfterExpiry(record)
return
}
if !claimed {
return
}
storeErrStreak := 0
for {
latest, getErr := store.GetLiveCall(context.Background(), record.CallHash)
if getErr != nil {
// 记录已被清理(closed TTL 到期)不是故障,直接退出。
if errors.Is(getErr, ErrLiveCallNotFound) {
return
}
// store 抖动不等于控制权被接管:有限次重试;仍失败则按
// record.ExpiresAt 兜底 finalize,保证 usage log 与租约释放不丢。
storeErrStreak++
if storeErrStreak >= liveObserverStoreRetryLimit {
s.finalizeLiveCallAfterExpiry(record)
return
}
time.Sleep(liveObserverStoreRetryInterval)
continue
}
storeErrStreak = 0
record = latest
if record.Controller != LiveControllerObserver {
return
}
if !time.Now().Before(record.ExpiresAt) {
s.finalizeLiveCall(record)
return
}
upstream, dialErr := s.dialLiveSideband(context.Background(), record)
if dialErr != nil {
if !s.waitForLiveObserverRetry(record) {
return
}
continue
}
runErr := s.runLiveObserverConnection(record, upstream)
_ = upstream.Close()
if errors.Is(runErr, ErrLiveControllerChanged) {
return
}
if liveSessionEnded(runErr) {
s.finalizeLiveCall(record)
return
}
if !s.waitForLiveObserverRetry(record) {
return
}
}
}
func (s *OpenAIGatewayService) runLiveObserverConnection(record *LiveCallRecord, upstream liveFrameConn) error {
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
frameCh := make(chan []byte, 1)
errCh := make(chan error, 1)
go func() {
for {
messageType, payload, err := upstream.ReadFrame(ctx)
if err != nil {
select {
case errCh <- liveSidebandReadError(err):
case <-ctx.Done():
}
return
}
if messageType == coderws.MessageText {
select {
case frameCh <- payload:
case <-ctx.Done():
return
}
}
}
}()
refreshTicker := time.NewTicker(liveLeaseRefreshInterval)
defer refreshTicker.Stop()
controllerTicker := time.NewTicker(liveObserverPollInterval)
defer controllerTicker.Stop()
maxTimer := time.NewTimer(time.Until(record.ExpiresAt))
defer maxTimer.Stop()
store, _ := s.liveStore()
for {
select {
case payload := <-frameCh:
eventType := strings.TrimSpace(gjson.GetBytes(payload, "type").String())
if eventType == "session.closed" || eventType == "session.ended" {
return ErrLiveCallNotFound
}
case err := <-errCh:
return err
case <-controllerTicker.C:
controller, err := store.GetLiveController(context.Background(), record.CallHash)
if err != nil {
return err
}
if controller != LiveControllerObserver {
return ErrLiveControllerChanged
}
case <-refreshTicker.C:
if !s.refreshLiveLease(record) {
return ErrLiveUnavailable
}
case <-maxTimer.C:
closeCtx, closeCancel := context.WithTimeout(context.Background(), 2*time.Second)
_ = upstream.WriteFrame(closeCtx, coderws.MessageText, []byte(`{"type":"session.close"}`))
closeCancel()
return context.DeadlineExceeded
}
}
}
func (s *OpenAIGatewayService) waitForLiveObserverRetry(record *LiveCallRecord) bool {
timer := time.NewTimer(time.Second)
defer timer.Stop()
<-timer.C
store, err := s.liveStore()
if err != nil {
return false
}
controller, getErr := store.GetLiveController(context.Background(), record.CallHash)
if getErr != nil && !errors.Is(getErr, ErrLiveCallNotFound) {
// store 报错不等于控制权被接管:返回 true 交回 observeLiveCall 循环顶部,
// 由它对 store 故障做有限次重试与 ExpiresAt 兜底 finalize。在这里返回
// false 会让 Redis 抖动时会话静默结束、不留记录。
return true
}
// 过期不在此处判定:返回 true 让调用方回到循环顶部的过期分支,由它 finalize
// (写 usage log + 释放租约)。在这里直接返回 false 会让会话静默结束、不留记录。
return getErr == nil && controller == LiveControllerObserver
}
// finalizeLiveCallAfterExpiry 是 store 持续报错、observer 无法继续观察时的兜底:
// 等到会话最长时限 ExpiresAt 再 finalize,保证 usage log 与租约释放最迟在会话到期
// 时完成。MarkLiveCallClosed 的 first 语义保证与其他恢复路径不会重复落库。
func (s *OpenAIGatewayService) finalizeLiveCallAfterExpiry(record *LiveCallRecord) {
if record == nil {
return
}
if wait := time.Until(record.ExpiresAt); wait > 0 {
time.Sleep(wait)
}
s.finalizeLiveCall(record)
}
func (s *OpenAIGatewayService) refreshLiveLease(record *LiveCallRecord) bool {
cache, err := s.liveConcurrencyCache()
if err != nil {
return false
}
ctx, cancel := context.WithTimeout(context.Background(), liveRedisOperationTimeout)
defer cancel()
refreshed, err := cache.RefreshLiveLease(ctx, record.AccountID, record.UserID, record.APIKeyID, record.LeaseID)
return err == nil && refreshed
}
func (s *OpenAIGatewayService) releaseLiveLease(accountID, userID, apiKeyID int64, leaseID string) {
cache, err := s.liveConcurrencyCache()
if err != nil {
return
}
ctx, cancel := context.WithTimeout(context.Background(), liveRedisOperationTimeout)
defer cancel()
_ = cache.ReleaseLiveLease(ctx, accountID, userID, apiKeyID, leaseID)
}
func (s *OpenAIGatewayService) finalizeLiveCall(record *LiveCallRecord) {
if record == nil {
return
}
store, err := s.liveStore()
if err != nil {
return
}
ctx, cancel := context.WithTimeout(context.Background(), liveRedisOperationTimeout)
first, err := store.MarkLiveCallClosed(ctx, record.CallHash, liveClosedRecordTTL)
cancel()
if err != nil || !first {
return
}
s.releaseLiveLease(record.AccountID, record.UserID, record.APIKeyID, record.LeaseID)
if s.usageLogRepo == nil {
return
}
duration := int(time.Since(record.CreatedAt).Milliseconds())
if duration < 0 {
duration = 0
}
inboundEndpoint := record.InboundEndpoint
upstreamEndpoint := "/backend-api/codex/realtime/calls"
userAgent := record.UserAgent
ipAddress := record.IPAddress
billingType := int8(BillingTypeBalance)
if record.SubscriptionID > 0 {
billingType = BillingTypeSubscription
}
// TODO(billing): Live 会话目前不计费:TotalCost/ActualCost 恒为 0,完全绕过
// recordUsageCore/applyUsageBilling,余额模式下极低余额也能反复开启最长
// liveMaxSessionDuration 的会话。若确认按时长计费,应在此接入计费管道;
// 若确认有意免费,删除本注释即可(零值行为由
// TestFinalizeLiveCallIsIdempotentAndWritesZeroUsage 锁定)。
//
// 这是该会话唯一一次落库机会(MarkLiveCallClosed 已标记 first),失败即永久
// 丢失,因此走带日志与同步兜底的 writeUsageLogBestEffortissue #3656)。
writeUsageLogBestEffort(context.Background(), s.usageLogRepo, &UsageLog{
UserID: record.UserID,
APIKeyID: record.APIKeyID,
AccountID: record.AccountID,
RequestID: record.CallHash,
Model: record.Model,
RequestedModel: record.Model,
GroupID: liveOptionalID(record.GroupID),
SubscriptionID: liveOptionalID(record.SubscriptionID),
RateMultiplier: 1,
BillingType: billingType,
RequestType: RequestTypeLive,
DurationMs: &duration,
UserAgent: &userAgent,
IPAddress: &ipAddress,
InboundEndpoint: &inboundEndpoint,
UpstreamEndpoint: &upstreamEndpoint,
CreatedAt: record.CreatedAt,
}, "service.openai_live")
}