From 1187390f9ea498f652e2f6d0815bfdf26c257b1a Mon Sep 17 00:00:00 2001 From: krakenalt Date: Thu, 20 Aug 2026 13:48:27 +0300 Subject: [PATCH 1/6] [feat]: add GigaChat provider core Register GigaChat and implement auth, inference, files, batches, and provider routing. --- core/bifrost.go | 3 + core/changelog.md | 6 + core/providers/gigachat/attachments_cache.go | 356 ++++ core/providers/gigachat/auth.go | 596 ++++++ core/providers/gigachat/batch.go | 708 +++++++ core/providers/gigachat/chat.go | 837 ++++++++ core/providers/gigachat/chat_attachments.go | 377 ++++ core/providers/gigachat/count_tokens.go | 208 ++ core/providers/gigachat/embedding.go | 170 ++ core/providers/gigachat/errors.go | 76 + core/providers/gigachat/files.go | 617 ++++++ core/providers/gigachat/gigachat.go | 1502 ++++++++++++++ core/providers/gigachat/models.go | 79 + core/providers/gigachat/responses.go | 1799 +++++++++++++++++ .../gigachat/responses_attachments.go | 304 +++ core/providers/gigachat/schema.go | 444 ++++ core/providers/gigachat/tools.go | 651 ++++++ core/providers/gigachat/types.go | 627 ++++++ core/providers/gigachat/utils.go | 449 ++++ core/schemas/account.go | 107 + core/schemas/bifrost.go | 2 + core/utils.go | 13 +- 22 files changed, 9929 insertions(+), 2 deletions(-) create mode 100644 core/providers/gigachat/attachments_cache.go create mode 100644 core/providers/gigachat/auth.go create mode 100644 core/providers/gigachat/batch.go create mode 100644 core/providers/gigachat/chat.go create mode 100644 core/providers/gigachat/chat_attachments.go create mode 100644 core/providers/gigachat/count_tokens.go create mode 100644 core/providers/gigachat/embedding.go create mode 100644 core/providers/gigachat/errors.go create mode 100644 core/providers/gigachat/files.go create mode 100644 core/providers/gigachat/gigachat.go create mode 100644 core/providers/gigachat/models.go create mode 100644 core/providers/gigachat/responses.go create mode 100644 core/providers/gigachat/responses_attachments.go create mode 100644 core/providers/gigachat/schema.go create mode 100644 core/providers/gigachat/tools.go create mode 100644 core/providers/gigachat/types.go create mode 100644 core/providers/gigachat/utils.go diff --git a/core/bifrost.go b/core/bifrost.go index a8623c4a4bd..1a8943cb55a 100644 --- a/core/bifrost.go +++ b/core/bifrost.go @@ -31,6 +31,7 @@ import ( "github.com/maximhq/bifrost/core/providers/elevenlabs" "github.com/maximhq/bifrost/core/providers/fireworks" "github.com/maximhq/bifrost/core/providers/gemini" + "github.com/maximhq/bifrost/core/providers/gigachat" "github.com/maximhq/bifrost/core/providers/groq" "github.com/maximhq/bifrost/core/providers/huggingface" "github.com/maximhq/bifrost/core/providers/mistral" @@ -4499,6 +4500,8 @@ func (bifrost *Bifrost) createBaseProvider(providerKey schemas.ModelProvider, co return wafer.NewWaferProvider(config, bifrost.logger) case schemas.Gemini: return gemini.NewGeminiProvider(config, bifrost.logger), nil + case schemas.GigaChat: + return gigachat.NewGigaChatProvider(config, bifrost.logger) case schemas.OpenRouter: return openrouter.NewOpenRouterProvider(config, bifrost.logger), nil case schemas.Elevenlabs: diff --git a/core/changelog.md b/core/changelog.md index d555caedf1a..7a56601ef8e 100644 --- a/core/changelog.md +++ b/core/changelog.md @@ -1,3 +1,9 @@ +- fix: apply GigaChat file-list limit and cursor pagination locally [@krakenalt](https://github.com/krakenalt) +- fix: guard GigaChat batch output downloads against an empty key set [@krakenalt](https://github.com/krakenalt) +- fix: keep GigaChat batch pagination provider-local without widening shared batch response schemas [@krakenalt](https://github.com/krakenalt) +- fix: finalize GigaChat Chat Completions and Responses streams across normal and large-response passthrough paths [@krakenalt](https://github.com/krakenalt) +- fix: cache GigaChat TLS clients without hot-path certificate file reads [@krakenalt](https://github.com/krakenalt) +- fix: tighten GigaChat attachment retries, auth cache cleanup, file-list validation, structured output handling, and batch key configuration [@krakenalt](https://github.com/krakenalt) - feat: support Gemini's server-side `toolCall`/`toolResponse` parts with `thoughtSignature` round-trip fidelity - server-side search rounds now surface as `web_search_call` items carrying their own call ID and queries, unmapped tool types are preserved on the native round-trip instead of being dropped, and each `thoughtSignature` appears exactly once across the reconstructed parts so Gemini accepts the replayed turn - feat: async 3D generation on Runware via `/videos` plus a raw `/runware_passthrough` route - `taskType` is now read from extra_params so any Runware async task can be driven through `/videos` (the 16:9 1080p width/height defaults now apply only to `videoInference`), `outputs.files[].url` is surfaced as `VideoOutput` URLs with the content type derived from the file extension, and the passthrough route forwards raw task arrays for capabilities with no first-class Bifrost surface such as upscaling and background removal - feat: surface Runware's provider-reported per-task `cost` across image, video/3D and passthrough so pricing uses the exact figure verbatim instead of a datasheet estimate - this matters for task types like 3D that have no datasheet rate; when no cost is reported the behavior is unchanged diff --git a/core/providers/gigachat/attachments_cache.go b/core/providers/gigachat/attachments_cache.go new file mode 100644 index 00000000000..726bf85850a --- /dev/null +++ b/core/providers/gigachat/attachments_cache.go @@ -0,0 +1,356 @@ +package gigachat + +import ( + "crypto/sha256" + "encoding/hex" + "runtime" + "strings" + "sync" + "time" + "weak" + + "github.com/google/uuid" + "github.com/maximhq/bifrost/core/schemas" +) + +const ( + gigaChatAttachmentCacheTTL = 15 * time.Minute + gigaChatAttachmentCacheSweepInterval = time.Minute +) + +type gigaChatAttachmentCacheContextKey struct{} + +var gigaChatAttachmentCacheKey = gigaChatAttachmentCacheContextKey{} + +// GigaChat uploads inline attachments before inference. The context only keeps +// a small cache ID; uploaded file metadata lives in this provider-owned manager. +type gigaChatAttachmentCacheManager struct { + mu sync.Mutex + entries map[string]*gigaChatAttachmentCacheEntry + sweepTimer *time.Timer +} + +type gigaChatAttachmentCacheEntry struct { + cache gigaChatAttachmentCache + expiresAt time.Time + writers int +} + +type gigaChatAttachmentCache struct { + mu sync.Mutex + chat map[gigaChatChatAttachmentCacheKey]gigaChatCachedAttachment[schemas.ChatContentBlock] + responses map[gigaChatResponsesAttachmentCacheKey]gigaChatCachedAttachment[schemas.ResponsesMessageContentBlock] +} + +type gigaChatCachedAttachment[T any] struct { + replacement T + expiresAt time.Time +} + +type gigaChatChatAttachmentCacheKey struct { + request weak.Pointer[schemas.BifrostChatRequest] + keyHash string + messageIndex int + blockIndex int +} + +type gigaChatResponsesAttachmentCacheKey struct { + request weak.Pointer[schemas.BifrostResponsesRequest] + keyHash string + messageIndex int + blockIndex int +} + +func newGigaChatAttachmentCacheManager() *gigaChatAttachmentCacheManager { + return &gigaChatAttachmentCacheManager{ + entries: make(map[string]*gigaChatAttachmentCacheEntry), + } +} + +func (manager *gigaChatAttachmentCacheManager) lookupCache(ctx *schemas.BifrostContext) *gigaChatAttachmentCache { + cache, _ := manager.cacheFor(ctx, false) + return cache +} + +func (manager *gigaChatAttachmentCacheManager) cacheForWrite(ctx *schemas.BifrostContext) (*gigaChatAttachmentCache, *gigaChatAttachmentCacheEntry) { + return manager.cacheFor(ctx, true) +} + +func (manager *gigaChatAttachmentCacheManager) cacheFor(ctx *schemas.BifrostContext, create bool) (*gigaChatAttachmentCache, *gigaChatAttachmentCacheEntry) { + if manager == nil || ctx == nil { + return nil, nil + } + + ctx = ctx.Root() + if ctx.Err() != nil { + return nil, nil + } + now := time.Now() + manager.mu.Lock() + cacheID, _ := ctx.Value(gigaChatAttachmentCacheKey).(string) + if cacheID == "" { + if !create { + manager.mu.Unlock() + return nil, nil + } + cacheID = uuid.NewString() + ctx.SetValue(gigaChatAttachmentCacheKey, cacheID) + } + + entry := manager.entries[cacheID] + if entry != nil && entry.writers == 0 && !entry.expiresAt.After(now) { + delete(manager.entries, cacheID) + entry = nil + } + if entry == nil { + if !create { + manager.mu.Unlock() + return nil, nil + } + entry = &gigaChatAttachmentCacheEntry{} + manager.entries[cacheID] = entry + } + entry.expiresAt = now.Add(gigaChatAttachmentCacheTTL) + if create { + entry.writers++ + } + manager.scheduleSweepLocked() + manager.mu.Unlock() + + return &entry.cache, entry +} + +func (manager *gigaChatAttachmentCacheManager) finishWrite(entry *gigaChatAttachmentCacheEntry) { + if manager == nil || entry == nil { + return + } + manager.mu.Lock() + defer manager.mu.Unlock() + if entry.writers > 0 { + entry.writers-- + } +} + +func (manager *gigaChatAttachmentCacheManager) scheduleSweepLocked() { + if manager.sweepTimer != nil || len(manager.entries) == 0 { + return + } + manager.sweepTimer = time.AfterFunc(gigaChatAttachmentCacheSweepInterval, manager.sweep) +} + +func (manager *gigaChatAttachmentCacheManager) sweep() { + manager.mu.Lock() + defer manager.mu.Unlock() + manager.sweepTimer = nil + manager.pruneEntriesLocked(time.Now()) + manager.scheduleSweepLocked() +} + +func (manager *gigaChatAttachmentCacheManager) pruneEntriesLocked(now time.Time) { + for cacheID, entry := range manager.entries { + if entry.writers > 0 { + continue + } + if !entry.expiresAt.After(now) { + delete(manager.entries, cacheID) + continue + } + if entry.cache.prune(now) { + delete(manager.entries, cacheID) + } + } +} + +func (cache *gigaChatAttachmentCache) prune(now time.Time) bool { + cache.mu.Lock() + defer cache.mu.Unlock() + for key, attachment := range cache.chat { + if key.request.Value() == nil || !attachment.expiresAt.After(now) { + delete(cache.chat, key) + } + } + for key, attachment := range cache.responses { + if key.request.Value() == nil || !attachment.expiresAt.After(now) { + delete(cache.responses, key) + } + } + return len(cache.chat) == 0 && len(cache.responses) == 0 +} + +func (provider *GigaChatProvider) getCachedGigaChatChatAttachment( + ctx *schemas.BifrostContext, + key schemas.Key, + request *schemas.BifrostChatRequest, + messageIndex int, + blockIndex int, +) (schemas.ChatContentBlock, bool) { + if request == nil { + return schemas.ChatContentBlock{}, false + } + cache := provider.attachmentCache.lookupCache(ctx) + if cache == nil { + return schemas.ChatContentBlock{}, false + } + + cacheKey := gigaChatChatAttachmentCacheKey{ + request: weak.Make(request), + keyHash: gigaChatAttachmentKeyHash(key), + messageIndex: messageIndex, + blockIndex: blockIndex, + } + cache.mu.Lock() + attachment, ok := cache.chat[cacheKey] + now := time.Now() + if ok && !attachment.expiresAt.After(now) { + delete(cache.chat, cacheKey) + attachment = gigaChatCachedAttachment[schemas.ChatContentBlock]{} + ok = false + } else if ok { + attachment.expiresAt = now.Add(gigaChatAttachmentCacheTTL) + cache.chat[cacheKey] = attachment + } + cache.mu.Unlock() + runtime.KeepAlive(request) + return attachment.replacement, ok +} + +func (provider *GigaChatProvider) setCachedGigaChatChatAttachment( + ctx *schemas.BifrostContext, + key schemas.Key, + request *schemas.BifrostChatRequest, + messageIndex int, + blockIndex int, + replacement schemas.ChatContentBlock, +) { + if request == nil { + return + } + cache, entry := provider.attachmentCache.cacheForWrite(ctx) + if cache == nil { + return + } + defer provider.attachmentCache.finishWrite(entry) + + cacheKey := gigaChatChatAttachmentCacheKey{ + request: weak.Make(request), + keyHash: gigaChatAttachmentKeyHash(key), + messageIndex: messageIndex, + blockIndex: blockIndex, + } + cache.mu.Lock() + if cache.chat == nil { + cache.chat = make(map[gigaChatChatAttachmentCacheKey]gigaChatCachedAttachment[schemas.ChatContentBlock]) + } + cache.chat[cacheKey] = gigaChatCachedAttachment[schemas.ChatContentBlock]{ + replacement: replacement, + expiresAt: time.Now().Add(gigaChatAttachmentCacheTTL), + } + cache.mu.Unlock() + runtime.KeepAlive(request) +} + +func (provider *GigaChatProvider) getCachedGigaChatResponsesAttachment( + ctx *schemas.BifrostContext, + key schemas.Key, + request *schemas.BifrostResponsesRequest, + messageIndex int, + blockIndex int, +) (schemas.ResponsesMessageContentBlock, bool) { + if request == nil { + return schemas.ResponsesMessageContentBlock{}, false + } + cache := provider.attachmentCache.lookupCache(ctx) + if cache == nil { + return schemas.ResponsesMessageContentBlock{}, false + } + + cacheKey := gigaChatResponsesAttachmentCacheKey{ + request: weak.Make(request), + keyHash: gigaChatAttachmentKeyHash(key), + messageIndex: messageIndex, + blockIndex: blockIndex, + } + cache.mu.Lock() + attachment, ok := cache.responses[cacheKey] + now := time.Now() + if ok && !attachment.expiresAt.After(now) { + delete(cache.responses, cacheKey) + attachment = gigaChatCachedAttachment[schemas.ResponsesMessageContentBlock]{} + ok = false + } else if ok { + attachment.expiresAt = now.Add(gigaChatAttachmentCacheTTL) + cache.responses[cacheKey] = attachment + } + cache.mu.Unlock() + runtime.KeepAlive(request) + return attachment.replacement, ok +} + +func (provider *GigaChatProvider) setCachedGigaChatResponsesAttachment( + ctx *schemas.BifrostContext, + key schemas.Key, + request *schemas.BifrostResponsesRequest, + messageIndex int, + blockIndex int, + replacement schemas.ResponsesMessageContentBlock, +) { + if request == nil { + return + } + cache, entry := provider.attachmentCache.cacheForWrite(ctx) + if cache == nil { + return + } + defer provider.attachmentCache.finishWrite(entry) + + cacheKey := gigaChatResponsesAttachmentCacheKey{ + request: weak.Make(request), + keyHash: gigaChatAttachmentKeyHash(key), + messageIndex: messageIndex, + blockIndex: blockIndex, + } + cache.mu.Lock() + if cache.responses == nil { + cache.responses = make(map[gigaChatResponsesAttachmentCacheKey]gigaChatCachedAttachment[schemas.ResponsesMessageContentBlock]) + } + cache.responses[cacheKey] = gigaChatCachedAttachment[schemas.ResponsesMessageContentBlock]{ + replacement: replacement, + expiresAt: time.Now().Add(gigaChatAttachmentCacheTTL), + } + cache.mu.Unlock() + runtime.KeepAlive(request) +} + +func gigaChatAttachmentKeyHash(key schemas.Key) string { + hash := sha256.New() + writePart := func(label string, value string) { + value = strings.TrimSpace(value) + if value == "" { + return + } + hash.Write([]byte(label)) + hash.Write([]byte{0}) + hash.Write([]byte(value)) + hash.Write([]byte{0}) + } + + writePart("id", key.ID) + writePart("name", key.Name) + writePart("value", key.Value.GetValue()) + + if key.GigaChatKeyConfig != nil { + config := key.GigaChatKeyConfig + writePart("credentials", config.Credentials.GetValue()) + writePart("scope", config.Scope) + writePart("user", config.User.GetValue()) + writePart("password", config.Password.GetValue()) + writePart("access_token", config.AccessToken.GetValue()) + writePart("auth_url", config.AuthURL) + writePart("base_url", config.BaseURL) + writePart("cert_file", config.CertFile) + writePart("key_file", config.KeyFile) + writePart("ca_bundle_file", config.CABundleFile) + } + + return hex.EncodeToString(hash.Sum(nil)) +} diff --git a/core/providers/gigachat/auth.go b/core/providers/gigachat/auth.go new file mode 100644 index 00000000000..c9d5a0fa9f1 --- /dev/null +++ b/core/providers/gigachat/auth.go @@ -0,0 +1,596 @@ +package gigachat + +import ( + "context" + "crypto/sha256" + "encoding/base64" + "encoding/hex" + "fmt" + "net/http" + "net/url" + "strings" + "sync" + "time" + + "github.com/bytedance/sonic" + "github.com/google/uuid" + + providerUtils "github.com/maximhq/bifrost/core/providers/utils" + schemas "github.com/maximhq/bifrost/core/schemas" + "github.com/valyala/fasthttp" +) + +const gigaChatOAuthRefreshLeeway = time.Minute + +const ( + gigaChatAuthorizationHeader = "Authorization" + gigaChatUserAgentHeader = "User-Agent" + gigaChatUserAgent = "GigaChat-Bifrost-Provider" +) + +var gigaChatContextHeaders = map[string]string{ + "x-session-id": "X-Session-ID", + "x-request-id": "X-Request-ID", + "x-service-id": "X-Service-ID", + "x-operation-id": "X-Operation-ID", + "x-client-id": "X-Client-ID", + "x-trace-id": "X-Trace-ID", + "x-agent-id": "X-Agent-ID", +} + +type gigaChatCachedToken struct { + accessToken string + expiresAt time.Time +} + +type gigaChatTokenCacheEntry struct { + mu sync.Mutex + token gigaChatCachedToken + // Guarded by gigaChatTokenCache.mu; prevents pruning entries while callers hold a pointer. + refCount int +} + +type gigaChatTokenCache struct { + mu sync.Mutex + entries map[string]*gigaChatTokenCacheEntry + now func() time.Time +} + +func newGigaChatTokenCache(now func() time.Time) *gigaChatTokenCache { + if now == nil { + now = time.Now + } + return &gigaChatTokenCache{ + entries: make(map[string]*gigaChatTokenCacheEntry), + now: now, + } +} + +func (provider *GigaChatProvider) buildAuthHeaders(ctx *schemas.BifrostContext, key schemas.Key) (map[string]string, *schemas.BifrostError) { + return provider.buildAuthHeadersWithRefresh(ctx, key, false) +} + +func (provider *GigaChatProvider) refreshAuthHeaders(ctx *schemas.BifrostContext, key schemas.Key) (map[string]string, *schemas.BifrostError) { + return provider.buildAuthHeadersWithRefresh(ctx, key, true) +} + +func (provider *GigaChatProvider) buildAuthHeadersWithRefresh(ctx *schemas.BifrostContext, key schemas.Key, forceRefresh bool) (map[string]string, *schemas.BifrostError) { + if bifrostErr := provider.rejectProviderAuthorizationExtraHeader(); bifrostErr != nil { + return nil, bifrostErr + } + if bifrostErr := rejectRequestAuthorizationExtraHeader(ctx); bifrostErr != nil { + return nil, bifrostErr + } + + headers := map[string]string{ + gigaChatUserAgentHeader: gigaChatUserAgent, + } + + accessToken, hasBearerAuth, bifrostErr := provider.resolveGigaChatAccessTokenWithRefresh(ctx, key, forceRefresh) + if bifrostErr != nil { + return nil, bifrostErr + } + if hasBearerAuth { + headers[gigaChatAuthorizationHeader] = "Bearer " + accessToken + } else if key.GigaChatKeyConfig == nil || !key.GigaChatKeyConfig.HasClientCertificateMaterial() { + return nil, newGigaChatConfigurationError("GigaChat authentication requires key.value access token, gigachat_key_config access_token, credentials, user/password auth material, or mTLS cert_file/key_file material") + } + + applyGigaChatProviderContextHeaders(headers, provider.networkConfig.ExtraHeaders) + applyGigaChatRequestContextHeaders(headers, ctx) + return headers, nil +} + +func (provider *GigaChatProvider) rejectProviderAuthorizationExtraHeader() *schemas.BifrostError { + if hasGigaChatHeader(provider.networkConfig.ExtraHeaders, gigaChatAuthorizationHeader) { + return newGigaChatConfigurationError("network_config.extra_headers cannot include Authorization for GigaChat; configure GigaChat auth material instead") + } + return nil +} + +func rejectRequestAuthorizationExtraHeader(ctx *schemas.BifrostContext) *schemas.BifrostError { + if _, ok := getGigaChatRequestExtraHeader(ctx, gigaChatAuthorizationHeader); ok { + return newGigaChatConfigurationError("request extra headers cannot include Authorization for GigaChat; configure GigaChat auth material instead") + } + return nil +} + +func hasGigaChatHeader(headers map[string]string, headerName string) bool { + for key := range headers { + if strings.EqualFold(strings.TrimSpace(key), headerName) { + return true + } + } + return false +} + +func applyGigaChatProviderContextHeaders(headers map[string]string, extraHeaders map[string]string) { + for key, value := range extraHeaders { + canonicalHeader, ok := getGigaChatContextHeaderName(key) + if !ok || canonicalHeader == gigaChatAuthorizationHeader { + continue + } + if strings.TrimSpace(value) != "" { + headers[canonicalHeader] = value + } + } +} + +func applyGigaChatRequestContextHeaders(headers map[string]string, ctx *schemas.BifrostContext) { + if ctx == nil { + return + } + extraHeaders, ok := ctx.Value(schemas.BifrostContextKeyExtraHeaders).(map[string][]string) + if !ok { + return + } + for key, values := range extraHeaders { + canonicalHeader, ok := getGigaChatContextHeaderName(key) + if !ok || canonicalHeader == gigaChatAuthorizationHeader { + continue + } + for _, value := range values { + if strings.TrimSpace(value) != "" { + headers[canonicalHeader] = value + break + } + } + } +} + +func getGigaChatRequestExtraHeader(ctx *schemas.BifrostContext, headerName string) (string, bool) { + if ctx == nil { + return "", false + } + extraHeaders, ok := ctx.Value(schemas.BifrostContextKeyExtraHeaders).(map[string][]string) + if !ok { + return "", false + } + for key, values := range extraHeaders { + if !strings.EqualFold(strings.TrimSpace(key), headerName) { + continue + } + for _, value := range values { + if strings.TrimSpace(value) != "" { + return value, true + } + } + } + return "", false +} + +func getGigaChatContextHeaderName(headerName string) (string, bool) { + canonicalHeader, ok := gigaChatContextHeaders[strings.ToLower(strings.TrimSpace(headerName))] + return canonicalHeader, ok +} + +func (provider *GigaChatProvider) getOAuthAccessToken(ctx *schemas.BifrostContext, key schemas.Key) (string, *schemas.BifrostError) { + return provider.getOAuthAccessTokenWithRefresh(ctx, key, false) +} + +func (provider *GigaChatProvider) getOAuthAccessTokenWithRefresh(ctx *schemas.BifrostContext, key schemas.Key, forceRefresh bool) (string, *schemas.BifrostError) { + authConfig, bifrostErr := resolveGigaChatOAuthConfig(key) + if bifrostErr != nil { + return "", bifrostErr + } + + cacheKey := buildGigaChatOAuthCacheKey(authConfig) + entry := provider.tokenCache.acquireEntry(cacheKey) + entry.mu.Lock() + defer provider.tokenCache.releaseEntry(cacheKey, entry) + defer entry.mu.Unlock() + + if !forceRefresh && entry.token.isValid(provider.tokenCache.now().Add(gigaChatOAuthRefreshLeeway)) { + return entry.token.accessToken, nil + } + + token, bifrostErr := provider.requestGigaChatOAuthToken(ctx, authConfig) + if bifrostErr != nil { + return "", bifrostErr + } + entry.token = token + return token.accessToken, nil +} + +func (provider *GigaChatProvider) getPasswordAccessToken(ctx *schemas.BifrostContext, key schemas.Key) (string, *schemas.BifrostError) { + return provider.getPasswordAccessTokenWithRefresh(ctx, key, false) +} + +func (provider *GigaChatProvider) getPasswordAccessTokenWithRefresh(ctx *schemas.BifrostContext, key schemas.Key, forceRefresh bool) (string, *schemas.BifrostError) { + authConfig, bifrostErr := provider.resolveGigaChatPasswordAuthConfig(key) + if bifrostErr != nil { + return "", bifrostErr + } + + cacheKey := buildGigaChatPasswordAuthCacheKey(authConfig) + entry := provider.tokenCache.acquireEntry(cacheKey) + entry.mu.Lock() + defer provider.tokenCache.releaseEntry(cacheKey, entry) + defer entry.mu.Unlock() + + if !forceRefresh && entry.token.isValid(provider.tokenCache.now().Add(gigaChatOAuthRefreshLeeway)) { + return entry.token.accessToken, nil + } + + token, bifrostErr := provider.requestGigaChatPasswordToken(ctx, authConfig) + if bifrostErr != nil { + return "", bifrostErr + } + entry.token = token + return token.accessToken, nil +} + +func (provider *GigaChatProvider) getGigaChatAccessToken(ctx *schemas.BifrostContext, key schemas.Key) (string, *schemas.BifrostError) { + return provider.getGigaChatAccessTokenWithRefresh(ctx, key, false) +} + +func (provider *GigaChatProvider) getGigaChatAccessTokenWithRefresh(ctx *schemas.BifrostContext, key schemas.Key, forceRefresh bool) (string, *schemas.BifrostError) { + accessToken, hasBearerAuth, bifrostErr := provider.resolveGigaChatAccessTokenWithRefresh(ctx, key, forceRefresh) + if bifrostErr != nil || hasBearerAuth { + return accessToken, bifrostErr + } + + return "", newGigaChatConfigurationError("GigaChat authentication requires key.value access token or gigachat_key_config access_token, credentials, or user/password auth material") +} + +func (provider *GigaChatProvider) resolveGigaChatAccessTokenWithRefresh(ctx *schemas.BifrostContext, key schemas.Key, forceRefresh bool) (string, bool, *schemas.BifrostError) { + keyConfig := key.GigaChatKeyConfig + if !forceRefresh { + if accessToken, isSet, bifrostErr := resolveGigaChatExplicitAccessToken(key); isSet || bifrostErr != nil { + return accessToken, isSet, bifrostErr + } + } + if keyConfig != nil && keyConfig.Credentials.IsSet() { + accessToken, bifrostErr := provider.getOAuthAccessTokenWithRefresh(ctx, key, forceRefresh) + return accessToken, true, bifrostErr + } + if keyConfig != nil && (keyConfig.User.IsSet() || keyConfig.Password.IsSet()) { + accessToken, bifrostErr := provider.getPasswordAccessTokenWithRefresh(ctx, key, forceRefresh) + return accessToken, true, bifrostErr + } + if accessToken, isSet, bifrostErr := resolveGigaChatExplicitAccessToken(key); isSet || bifrostErr != nil { + return accessToken, isSet, bifrostErr + } + + return "", false, nil +} + +func resolveGigaChatExplicitAccessToken(key schemas.Key) (string, bool, *schemas.BifrostError) { + if key.GigaChatKeyConfig != nil && key.GigaChatKeyConfig.AccessToken.IsSet() { + accessToken := strings.TrimSpace(key.GigaChatKeyConfig.AccessToken.GetValue()) + if accessToken == "" { + return "", true, newGigaChatConfigurationError("gigachat_key_config.access_token resolved to an empty value") + } + return accessToken, true, nil + } + if key.Value.IsSet() { + accessToken := strings.TrimSpace(key.Value.GetValue()) + if accessToken == "" { + return "", true, newGigaChatConfigurationError("GigaChat key value resolved to an empty access token") + } + return accessToken, true, nil + } + return "", false, nil +} + +func (cache *gigaChatTokenCache) acquireEntry(cacheKey string) *gigaChatTokenCacheEntry { + cache.mu.Lock() + defer cache.mu.Unlock() + + cache.pruneExpiredEntriesLocked(cache.now()) + + entry := cache.entries[cacheKey] + if entry == nil { + entry = &gigaChatTokenCacheEntry{} + cache.entries[cacheKey] = entry + } + entry.refCount++ + return entry +} + +func (cache *gigaChatTokenCache) releaseEntry(cacheKey string, entry *gigaChatTokenCacheEntry) { + cache.mu.Lock() + defer cache.mu.Unlock() + + if entry.refCount > 0 { + entry.refCount-- + } + if entry.refCount != 0 || cache.entries[cacheKey] != entry { + return + } + + entry.mu.Lock() + reusable := entry.token.isValid(cache.now()) + entry.mu.Unlock() + if !reusable { + delete(cache.entries, cacheKey) + } +} + +func (cache *gigaChatTokenCache) pruneExpiredEntriesLocked(now time.Time) { + for cacheKey, entry := range cache.entries { + if entry.refCount != 0 { + continue + } + entry.mu.Lock() + reusable := entry.token.isValid(now) + entry.mu.Unlock() + + if !reusable { + delete(cache.entries, cacheKey) + } + } +} + +func (token gigaChatCachedToken) isValid(validAfter time.Time) bool { + return token.accessToken != "" && token.expiresAt.After(validAfter) +} + +func parseGigaChatExpiresAt(value int64) time.Time { + if value > 1_000_000_000_000 { + return time.UnixMilli(value) + } + return time.Unix(value, 0) +} + +type gigaChatOAuthConfig struct { + authURL string + credentials string + scope string + keyConfig *schemas.GigaChatKeyConfig +} + +type gigaChatPasswordAuthConfig struct { + tokenURL string + user string + password string + keyConfig *schemas.GigaChatKeyConfig +} + +func resolveGigaChatOAuthConfig(key schemas.Key) (gigaChatOAuthConfig, *schemas.BifrostError) { + keyConfig := key.GigaChatKeyConfig + if keyConfig == nil || !keyConfig.Credentials.IsSet() { + return gigaChatOAuthConfig{}, newGigaChatConfigurationError("gigachat_key_config.credentials is required for OAuth token exchange") + } + + credentials := strings.TrimSpace(keyConfig.Credentials.GetValue()) + if credentials == "" { + return gigaChatOAuthConfig{}, newGigaChatConfigurationError("gigachat_key_config.credentials resolved to an empty value") + } + + scope := strings.TrimSpace(keyConfig.Scope) + if scope == "" { + scope = schemas.DefaultGigaChatScope + } + + return gigaChatOAuthConfig{ + authURL: resolveAuthURL(key), + credentials: credentials, + scope: scope, + keyConfig: keyConfig, + }, nil +} + +func buildGigaChatOAuthCacheKey(authConfig gigaChatOAuthConfig) string { + tlsFingerprint := gigaChatAuthTLSConfigFingerprint(authConfig.keyConfig) + hash := sha256.New() + hash.Write([]byte("oauth")) + hash.Write([]byte{0}) + hash.Write([]byte(authConfig.authURL)) + hash.Write([]byte{0}) + hash.Write([]byte(authConfig.scope)) + hash.Write([]byte{0}) + hash.Write([]byte(authConfig.credentials)) + hash.Write([]byte{0}) + hash.Write([]byte(tlsFingerprint)) + return hex.EncodeToString(hash.Sum(nil)) +} + +func (provider *GigaChatProvider) resolveGigaChatPasswordAuthConfig(key schemas.Key) (gigaChatPasswordAuthConfig, *schemas.BifrostError) { + keyConfig := key.GigaChatKeyConfig + if keyConfig == nil || !keyConfig.User.IsSet() || !keyConfig.Password.IsSet() { + return gigaChatPasswordAuthConfig{}, newGigaChatConfigurationError("gigachat_key_config.user and gigachat_key_config.password are required for password auth") + } + + user := keyConfig.User.GetValue() + if strings.TrimSpace(user) == "" { + return gigaChatPasswordAuthConfig{}, newGigaChatConfigurationError("gigachat_key_config.user resolved to an empty value") + } + password := keyConfig.Password.GetValue() + if strings.TrimSpace(password) == "" { + return gigaChatPasswordAuthConfig{}, newGigaChatConfigurationError("gigachat_key_config.password resolved to an empty value") + } + + baseURL := resolveBaseURL(key, provider.networkConfig) + return gigaChatPasswordAuthConfig{ + tokenURL: buildGigaChatURL(baseURL, gigaChatAPIVersionV1, "/token"), + user: user, + password: password, + keyConfig: keyConfig, + }, nil +} + +func buildGigaChatPasswordAuthCacheKey(authConfig gigaChatPasswordAuthConfig) string { + tlsFingerprint := gigaChatAuthTLSConfigFingerprint(authConfig.keyConfig) + hash := sha256.New() + hash.Write([]byte("password")) + hash.Write([]byte{0}) + hash.Write([]byte(authConfig.tokenURL)) + hash.Write([]byte{0}) + hash.Write([]byte(authConfig.user)) + hash.Write([]byte{0}) + hash.Write([]byte(authConfig.password)) + hash.Write([]byte{0}) + hash.Write([]byte(tlsFingerprint)) + return hex.EncodeToString(hash.Sum(nil)) +} + +func (provider *GigaChatProvider) requestGigaChatOAuthToken(ctx *schemas.BifrostContext, authConfig gigaChatOAuthConfig) (gigaChatCachedToken, *schemas.BifrostError) { + if ctx == nil { + ctx = schemas.NewBifrostContext(context.Background(), schemas.NoDeadline) + } + + req := fasthttp.AcquireRequest() + resp := fasthttp.AcquireResponse() + defer fasthttp.ReleaseRequest(req) + defer fasthttp.ReleaseResponse(resp) + + form := url.Values{} + form.Set("scope", authConfig.scope) + + req.SetRequestURI(authConfig.authURL) + req.Header.SetMethod(http.MethodPost) + req.Header.SetContentType("application/x-www-form-urlencoded") + req.Header.Set("Accept", "application/json") + req.Header.Set("RqUID", uuid.NewString()) + req.Header.Set(gigaChatUserAgentHeader, gigaChatUserAgent) + req.Header.Set("Authorization", "Basic "+authConfig.credentials) + req.SetBodyString(form.Encode()) + + client, err := provider.getGigaChatTLSClient(provider.client, gigaChatTLSClientCacheAuth, gigaChatAuthTLSKeyConfig(authConfig.keyConfig)) + if err != nil { + return gigaChatCachedToken{}, newGigaChatConfigurationError(err.Error()) + } + + _, bifrostErr, wait := providerUtils.MakeRequestWithContext(ctx, client, req, resp) + defer wait() + if bifrostErr != nil { + bifrostErr.ExtraFields.Provider = provider.GetProviderKey() + return gigaChatCachedToken{}, bifrostErr + } + + if resp.StatusCode() < http.StatusOK || resp.StatusCode() >= http.StatusMultipleChoices { + return gigaChatCachedToken{}, ParseGigaChatError(resp, provider.GetProviderKey()) + } + + body, err := providerUtils.CheckAndDecodeBody(resp) + if err != nil { + return gigaChatCachedToken{}, newGigaChatProviderResponseError("failed to decode GigaChat token response", err) + } + + var tokenResponse GigaChatTokenResponse + if err := sonic.Unmarshal(body, &tokenResponse); err != nil { + return gigaChatCachedToken{}, newGigaChatProviderResponseError("failed to parse GigaChat token response", err) + } + if strings.TrimSpace(tokenResponse.AccessToken) == "" { + return gigaChatCachedToken{}, newGigaChatProviderResponseError("GigaChat token response missing access_token", nil) + } + if tokenResponse.ExpiresAt <= 0 { + return gigaChatCachedToken{}, newGigaChatProviderResponseError("GigaChat token response missing expires_at", nil) + } + + expiresAt := parseGigaChatExpiresAt(tokenResponse.ExpiresAt) + if !expiresAt.After(provider.tokenCache.now()) { + return gigaChatCachedToken{}, newGigaChatProviderResponseError("GigaChat token response is already expired", nil) + } + + return gigaChatCachedToken{ + accessToken: tokenResponse.AccessToken, + expiresAt: expiresAt, + }, nil +} + +func (provider *GigaChatProvider) requestGigaChatPasswordToken(ctx *schemas.BifrostContext, authConfig gigaChatPasswordAuthConfig) (gigaChatCachedToken, *schemas.BifrostError) { + if ctx == nil { + ctx = schemas.NewBifrostContext(context.Background(), schemas.NoDeadline) + } + + req := fasthttp.AcquireRequest() + resp := fasthttp.AcquireResponse() + defer fasthttp.ReleaseRequest(req) + defer fasthttp.ReleaseResponse(resp) + + req.SetRequestURI(authConfig.tokenURL) + req.Header.SetMethod(http.MethodPost) + req.Header.SetContentType("application/x-www-form-urlencoded") + req.Header.Set("Accept", "application/json") + req.Header.Set(gigaChatUserAgentHeader, gigaChatUserAgent) + req.Header.Set("Authorization", "Basic "+base64.StdEncoding.EncodeToString([]byte(authConfig.user+":"+authConfig.password))) + + client, err := provider.getGigaChatTLSClient(provider.client, gigaChatTLSClientCacheAuth, gigaChatAuthTLSKeyConfig(authConfig.keyConfig)) + if err != nil { + return gigaChatCachedToken{}, newGigaChatConfigurationError(err.Error()) + } + + _, bifrostErr, wait := providerUtils.MakeRequestWithContext(ctx, client, req, resp) + defer wait() + if bifrostErr != nil { + bifrostErr.ExtraFields.Provider = provider.GetProviderKey() + return gigaChatCachedToken{}, bifrostErr + } + + if resp.StatusCode() < http.StatusOK || resp.StatusCode() >= http.StatusMultipleChoices { + return gigaChatCachedToken{}, ParseGigaChatError(resp, provider.GetProviderKey()) + } + + body, err := providerUtils.CheckAndDecodeBody(resp) + if err != nil { + return gigaChatCachedToken{}, newGigaChatProviderResponseError("failed to decode GigaChat password token response", err) + } + + var tokenResponse GigaChatPasswordTokenResponse + if err := sonic.Unmarshal(body, &tokenResponse); err != nil { + return gigaChatCachedToken{}, newGigaChatProviderResponseError("failed to parse GigaChat password token response", err) + } + if strings.TrimSpace(tokenResponse.Token) == "" { + return gigaChatCachedToken{}, newGigaChatProviderResponseError("GigaChat password token response missing tok", nil) + } + if tokenResponse.ExpiresAt <= 0 { + return gigaChatCachedToken{}, newGigaChatProviderResponseError("GigaChat password token response missing exp", nil) + } + + expiresAt := parseGigaChatExpiresAt(tokenResponse.ExpiresAt) + if !expiresAt.After(provider.tokenCache.now()) { + return gigaChatCachedToken{}, newGigaChatProviderResponseError("GigaChat password token response is already expired", nil) + } + + return gigaChatCachedToken{ + accessToken: tokenResponse.Token, + expiresAt: expiresAt, + }, nil +} + +func newGigaChatConfigurationError(message string) *schemas.BifrostError { + bifrostErr := providerUtils.NewConfigurationError(message) + bifrostErr.ExtraFields.Provider = schemas.GigaChat + return bifrostErr +} + +func newGigaChatProviderResponseError(message string, err error) *schemas.BifrostError { + statusCode := http.StatusBadGateway + bifrostErr := &schemas.BifrostError{ + IsBifrostError: false, + StatusCode: &statusCode, + Error: &schemas.ErrorField{ + Message: message, + Error: err, + }, + ExtraFields: schemas.BifrostErrorExtraFields{ + Provider: schemas.GigaChat, + }, + } + if err != nil { + bifrostErr.Error.Message = fmt.Sprintf("%s: %v", message, err) + } + return bifrostErr +} diff --git a/core/providers/gigachat/batch.go b/core/providers/gigachat/batch.go new file mode 100644 index 00000000000..3bacbcb6ace --- /dev/null +++ b/core/providers/gigachat/batch.go @@ -0,0 +1,708 @@ +package gigachat + +import ( + "bytes" + "encoding/json" + "fmt" + "net/http" + "net/url" + "strconv" + "strings" + "time" + + openaiProvider "github.com/maximhq/bifrost/core/providers/openai" + providerUtils "github.com/maximhq/bifrost/core/providers/utils" + schemas "github.com/maximhq/bifrost/core/schemas" +) + +type openAICompatibleBatchInputRow struct { + CustomID string `json:"custom_id"` + Method string `json:"method,omitempty"` + URL string `json:"url,omitempty"` + Body json.RawMessage `json:"body,omitempty"` +} + +const gigaChatBatchCompletionWindow24h = "24h" + +func (provider *GigaChatProvider) batchCreateWithRefresh(ctx *schemas.BifrostContext, key schemas.Key, request *schemas.BifrostBatchCreateRequest, forceRefresh bool) (*schemas.BifrostBatchCreateResponse, *schemas.BifrostError) { + ctx = ensureGigaChatContext(ctx) + + method, completionWindow, body, bifrostErr := provider.buildGigaChatBatchCreatePayload(ctx, key, request) + if bifrostErr != nil { + return nil, bifrostErr + } + + values := url.Values{} + values.Set("method", string(method)) + path := withGigaChatQuery("/batches", values) + responseBody, providerResponseHeaders, _, latency, bifrostErr := provider.executeGigaChatFileRequest(ctx, key, schemas.BatchCreateRequest, http.MethodPost, path, "application/octet-stream", "application/json", body, body, forceRefresh) + if bifrostErr != nil { + return nil, bifrostErr + } + + sendBackRawRequest := providerUtils.ShouldSendBackRawRequest(ctx, provider.sendBackRawRequest) + sendBackRawResponse := providerUtils.ShouldSendBackRawResponse(ctx, provider.sendBackRawResponse) + + var raw json.RawMessage + rawRequest, rawResponse, bifrostErr := providerUtils.HandleProviderResponse(responseBody, &raw, body, sendBackRawRequest, sendBackRawResponse) + if bifrostErr != nil { + return nil, enrichGigaChatError(ctx, bifrostErr, body, responseBody, sendBackRawRequest, sendBackRawResponse) + } + + batch, err := decodeGigaChatBatchResponse(responseBody) + if err != nil { + return nil, newGigaChatProviderResponseError("failed to decode GigaChat batch create response", err) + } + + response := toBifrostGigaChatBatchCreateResponse(provider.GetProviderKey(), batch, request, completionWindow, latency) + response.ExtraFields.ProviderResponseHeaders = providerResponseHeaders + if sendBackRawRequest { + response.ExtraFields.RawRequest = rawRequest + } + if sendBackRawResponse { + response.ExtraFields.RawResponse = rawResponse + } + return response, nil +} + +func (provider *GigaChatProvider) batchListWithRefresh(ctx *schemas.BifrostContext, key schemas.Key, _ *schemas.BifrostBatchListRequest, forceRefresh bool) (*schemas.BifrostBatchListResponse, *schemas.BifrostError) { + ctx = ensureGigaChatContext(ctx) + + responseBody, providerResponseHeaders, _, latency, bifrostErr := provider.executeGigaChatFileRequest(ctx, key, schemas.BatchListRequest, http.MethodGet, "/batches", "", "application/json", nil, nil, forceRefresh) + if bifrostErr != nil { + return nil, bifrostErr + } + + sendBackRawRequest := providerUtils.ShouldSendBackRawRequest(ctx, provider.sendBackRawRequest) + sendBackRawResponse := providerUtils.ShouldSendBackRawResponse(ctx, provider.sendBackRawResponse) + + gigaChatResponse := &GigaChatBatches{} + rawRequest, rawResponse, bifrostErr := providerUtils.HandleProviderResponse(responseBody, gigaChatResponse, nil, sendBackRawRequest, sendBackRawResponse) + if bifrostErr != nil { + return nil, enrichGigaChatError(ctx, bifrostErr, nil, responseBody, sendBackRawRequest, sendBackRawResponse) + } + + response := toBifrostGigaChatBatchListResponse(provider.GetProviderKey(), *gigaChatResponse, latency) + response.ExtraFields.ProviderResponseHeaders = providerResponseHeaders + if sendBackRawRequest { + response.ExtraFields.RawRequest = rawRequest + } + if sendBackRawResponse { + response.ExtraFields.RawResponse = rawResponse + } + return response, nil +} + +func (provider *GigaChatProvider) batchRetrieveWithRefresh(ctx *schemas.BifrostContext, key schemas.Key, request *schemas.BifrostBatchRetrieveRequest, forceRefresh bool) (*schemas.BifrostBatchRetrieveResponse, *schemas.BifrostError) { + ctx = ensureGigaChatContext(ctx) + + values := url.Values{} + values.Set("batch_id", strings.TrimSpace(request.BatchID)) + path := withGigaChatQuery("/batches", values) + responseBody, providerResponseHeaders, _, latency, bifrostErr := provider.executeGigaChatFileRequest(ctx, key, schemas.BatchRetrieveRequest, http.MethodGet, path, "", "application/json", nil, nil, forceRefresh) + if bifrostErr != nil { + return nil, bifrostErr + } + + sendBackRawRequest := providerUtils.ShouldSendBackRawRequest(ctx, provider.sendBackRawRequest) + sendBackRawResponse := providerUtils.ShouldSendBackRawResponse(ctx, provider.sendBackRawResponse) + + var raw json.RawMessage + rawRequest, rawResponse, bifrostErr := providerUtils.HandleProviderResponse(responseBody, &raw, nil, sendBackRawRequest, sendBackRawResponse) + if bifrostErr != nil { + return nil, enrichGigaChatError(ctx, bifrostErr, nil, responseBody, sendBackRawRequest, sendBackRawResponse) + } + + batch, err := decodeGigaChatBatchResponse(responseBody) + if err != nil { + return nil, newGigaChatProviderResponseError("failed to decode GigaChat batch retrieve response", err) + } + + response := toBifrostGigaChatBatchRetrieveResponse(provider.GetProviderKey(), batch, "", latency) + response.ExtraFields.ProviderResponseHeaders = providerResponseHeaders + if sendBackRawRequest { + response.ExtraFields.RawRequest = rawRequest + } + if sendBackRawResponse { + response.ExtraFields.RawResponse = rawResponse + } + return response, nil +} + +func (provider *GigaChatProvider) buildGigaChatBatchCreatePayload(ctx *schemas.BifrostContext, key schemas.Key, request *schemas.BifrostBatchCreateRequest) (GigaChatBatchMethod, string, []byte, *schemas.BifrostError) { + method, completionWindow, bifrostErr := validateGigaChatBatchCreateRequest(request) + if bifrostErr != nil { + return "", "", nil, bifrostErr + } + + var ( + body []byte + err error + ) + switch { + case strings.TrimSpace(request.InputFileID) != "": + content, contentErr := provider.readGigaChatBatchInputFile(ctx, key, request) + if contentErr != nil { + return "", "", nil, contentErr + } + body, err = convertGigaChatBatchInputJSONL(request.Endpoint, content) + case len(request.Requests) > 0: + body, err = convertGigaChatBatchRequestItemsToJSONL(request.Endpoint, request.Requests) + default: + return "", "", nil, providerUtils.NewBifrostOperationError("either input_file_id or requests array is required for GigaChat batch API", nil) + } + if err != nil { + return "", "", nil, providerUtils.NewBifrostOperationError("failed to convert GigaChat batch input rows", err) + } + return method, completionWindow, body, nil +} + +func (provider *GigaChatProvider) readGigaChatBatchInputFile(ctx *schemas.BifrostContext, key schemas.Key, request *schemas.BifrostBatchCreateRequest) ([]byte, *schemas.BifrostError) { + fileRequest := &schemas.BifrostFileContentRequest{ + Provider: provider.GetProviderKey(), + Model: request.Model, + FileID: strings.TrimSpace(request.InputFileID), + } + response, bifrostErr := provider.fileContentWithRefresh(ctx, key, fileRequest, false) + if isGigaChatUnauthorizedError(bifrostErr) { + response, bifrostErr = provider.fileContentWithRefresh(ctx, key, fileRequest, true) + } + if bifrostErr != nil { + return nil, bifrostErr + } + return response.Content, nil +} + +func validateGigaChatBatchCreateRequest(request *schemas.BifrostBatchCreateRequest) (GigaChatBatchMethod, string, *schemas.BifrostError) { + if request == nil { + return "", "", providerUtils.NewBifrostOperationError("batch create request is nil", nil) + } + if strings.TrimSpace(string(request.Endpoint)) == "" { + return "", "", providerUtils.NewBifrostOperationError("endpoint is required for GigaChat batch API", nil) + } + method, err := toGigaChatBatchMethod(request.Endpoint) + if err != nil { + return "", "", providerUtils.NewBifrostOperationError(err.Error(), err) + } + if strings.TrimSpace(request.InputFileID) != "" && len(request.Requests) > 0 { + return "", "", providerUtils.NewBifrostOperationError("input_file_id and requests array cannot both be set for GigaChat batch API", nil) + } + completionWindow := strings.TrimSpace(request.CompletionWindow) + if completionWindow == "" { + completionWindow = gigaChatBatchCompletionWindow24h + } + if completionWindow != gigaChatBatchCompletionWindow24h { + return "", "", providerUtils.NewBifrostOperationError("GigaChat batch API supports completion_window=24h only", nil) + } + if request.InputBlob != nil { + return "", "", providerUtils.NewBifrostOperationError("GigaChat batch API does not support input_blob", nil) + } + if request.OutputFolder != nil { + return "", "", providerUtils.NewBifrostOperationError("GigaChat batch API does not support output_folder", nil) + } + if request.OutputExpiresAfter != nil { + return "", "", providerUtils.NewBifrostOperationError("GigaChat batch API does not support output_expires_after", nil) + } + if len(request.Metadata) > 0 { + return "", "", providerUtils.NewBifrostOperationError("GigaChat batch API does not support metadata", nil) + } + if len(request.ExtraParams) > 0 { + return "", "", providerUtils.NewBifrostOperationError("GigaChat batch API does not support extra batch create parameters", nil) + } + return method, completionWindow, nil +} + +func validateGigaChatBatchListRequest(request *schemas.BifrostBatchListRequest) *schemas.BifrostError { + if request.BeforeID != nil && strings.TrimSpace(*request.BeforeID) != "" { + return providerUtils.NewBifrostOperationError("GigaChat batch list does not support before_id pagination", nil) + } + if request.AfterID != nil && strings.TrimSpace(*request.AfterID) != "" { + return providerUtils.NewBifrostOperationError("GigaChat batch list does not support after_id pagination", nil) + } + if request.PageToken != nil && strings.TrimSpace(*request.PageToken) != "" { + return providerUtils.NewBifrostOperationError("GigaChat batch list does not support page_token pagination", nil) + } + if request.PageSize > 0 { + return providerUtils.NewBifrostOperationError("GigaChat batch list does not support page_size pagination", nil) + } + if request.NextCursor != nil && strings.TrimSpace(*request.NextCursor) != "" { + return providerUtils.NewBifrostOperationError("GigaChat batch list does not support next_cursor pagination", nil) + } + if len(request.ExtraParams) > 0 { + return providerUtils.NewBifrostOperationError("GigaChat batch list does not support extra parameters", nil) + } + return nil +} + +func toGigaChatBatchMethod(endpoint schemas.BatchEndpoint) (GigaChatBatchMethod, error) { + switch normalizeGigaChatBatchEndpoint(string(endpoint)) { + case "/v1/chat/completions", "/chat/completions": + return GigaChatBatchMethodChatCompletions, nil + case "/v1/responses", "/responses": + return GigaChatBatchMethodResponses, nil + case "/v1/embeddings", "/embeddings": + return GigaChatBatchMethodEmbedder, nil + default: + return "", fmt.Errorf("GigaChat batches do not support endpoint %q", endpoint) + } +} + +func normalizeGigaChatBatchEndpoint(endpoint string) string { + normalized := strings.ToLower(strings.TrimSpace(endpoint)) + if normalized == "" { + return "" + } + if !strings.HasPrefix(normalized, "/") { + normalized = "/" + normalized + } + return strings.TrimRight(normalized, "/") +} + +func toBifrostGigaChatBatchStatus(status GigaChatBatchStatus) schemas.BatchStatus { + switch status { + case GigaChatBatchStatusCreated: + return schemas.BatchStatusValidating + case GigaChatBatchStatusInProgress: + return schemas.BatchStatusInProgress + case GigaChatBatchStatusCompleted: + return schemas.BatchStatusCompleted + default: + return schemas.BatchStatus(status) + } +} + +func decodeGigaChatBatchResponse(responseBody []byte) (GigaChatBatch, error) { + var batch GigaChatBatch + if err := json.Unmarshal(responseBody, &batch); err == nil && strings.TrimSpace(batch.ID) != "" { + return batch, nil + } + + var batches GigaChatBatches + if err := json.Unmarshal(responseBody, &batches); err != nil { + return GigaChatBatch{}, err + } + switch len(batches.Data) { + case 0: + return GigaChatBatch{}, fmt.Errorf("GigaChat batch response does not contain batch data") + case 1: + if strings.TrimSpace(batches.Data[0].ID) == "" { + return GigaChatBatch{}, fmt.Errorf("GigaChat batch response is missing id") + } + return batches.Data[0], nil + default: + return GigaChatBatch{}, fmt.Errorf("GigaChat batch response contains %d batches, want 1", len(batches.Data)) + } +} + +func toBifrostGigaChatBatchCreateResponse(providerName schemas.ModelProvider, batch GigaChatBatch, request *schemas.BifrostBatchCreateRequest, completionWindow string, latency time.Duration) *schemas.BifrostBatchCreateResponse { + inputFileID := "" + endpoint := "" + if request != nil { + inputFileID = strings.TrimSpace(request.InputFileID) + endpoint = string(request.Endpoint) + } + if batch.InputFileID != nil && strings.TrimSpace(*batch.InputFileID) != "" { + inputFileID = strings.TrimSpace(*batch.InputFileID) + } + if strings.TrimSpace(batch.CompletionWindow) != "" { + completionWindow = batch.CompletionWindow + } + if strings.TrimSpace(endpoint) == "" { + endpoint = toBifrostGigaChatBatchEndpoint(batch.Method) + } + + response := &schemas.BifrostBatchCreateResponse{ + ID: batch.ID, + Object: toBifrostGigaChatBatchObject(batch.Object), + Endpoint: endpoint, + InputFileID: inputFileID, + CompletionWindow: completionWindow, + Status: toBifrostGigaChatBatchStatus(batch.Status), + RequestCounts: toBifrostGigaChatBatchRequestCounts(batch.RequestCounts), + CreatedAt: batch.CreatedAt, + OutputFileID: toBifrostGigaChatBatchOutputFileID(batch), + ErrorFileID: cleanGigaChatBatchFileID(batch.ErrorFileID), + ExtraFields: schemas.BifrostResponseExtraFields{ + Provider: providerName, + Latency: latency.Milliseconds(), + }, + } + return response +} + +func toBifrostGigaChatBatchListResponse(providerName schemas.ModelProvider, batches GigaChatBatches, latency time.Duration) *schemas.BifrostBatchListResponse { + data := make([]schemas.BifrostBatchRetrieveResponse, 0, len(batches.Data)) + for _, batch := range batches.Data { + data = append(data, *toBifrostGigaChatBatchRetrieveResponse(providerName, batch, "", latency)) + } + + response := &schemas.BifrostBatchListResponse{ + Object: "list", + Data: data, + ExtraFields: schemas.BifrostResponseExtraFields{ + Provider: providerName, + Latency: latency.Milliseconds(), + }, + } + if len(data) > 0 { + firstID := data[0].ID + lastID := data[len(data)-1].ID + response.FirstID = &firstID + response.LastID = &lastID + } + return response +} + +func toBifrostGigaChatBatchRetrieveResponse(providerName schemas.ModelProvider, batch GigaChatBatch, fallbackEndpoint string, latency time.Duration) *schemas.BifrostBatchRetrieveResponse { + inputFileID := "" + if batch.InputFileID != nil { + inputFileID = strings.TrimSpace(*batch.InputFileID) + } + endpoint := fallbackEndpoint + if strings.TrimSpace(endpoint) == "" { + endpoint = toBifrostGigaChatBatchEndpoint(batch.Method) + } + + return &schemas.BifrostBatchRetrieveResponse{ + ID: batch.ID, + Object: toBifrostGigaChatBatchObject(batch.Object), + Endpoint: endpoint, + InputFileID: inputFileID, + CompletionWindow: batch.CompletionWindow, + Status: toBifrostGigaChatBatchStatus(batch.Status), + RequestCounts: toBifrostGigaChatBatchRequestCounts(batch.RequestCounts), + CreatedAt: batch.CreatedAt, + CompletedAt: batch.CompletedAt, + OutputFileID: toBifrostGigaChatBatchOutputFileID(batch), + ErrorFileID: cleanGigaChatBatchFileID(batch.ErrorFileID), + ExtraFields: schemas.BifrostResponseExtraFields{ + Provider: providerName, + Latency: latency.Milliseconds(), + }, + } +} + +func toBifrostGigaChatBatchRequestCounts(counts *GigaChatBatchRequestCounts) schemas.BatchRequestCounts { + if counts == nil { + return schemas.BatchRequestCounts{} + } + return schemas.BatchRequestCounts{ + Total: counts.Total, + Completed: counts.Completed, + Failed: counts.Failed, + } +} + +func toBifrostGigaChatBatchEndpoint(method GigaChatBatchMethod) string { + switch method { + case GigaChatBatchMethodChatCompletions: + return string(schemas.BatchEndpointChatCompletions) + case GigaChatBatchMethodResponses: + return string(schemas.BatchEndpointResponses) + case GigaChatBatchMethodEmbedder: + return string(schemas.BatchEndpointEmbeddings) + default: + return "" + } +} + +func toBifrostGigaChatBatchObject(object string) string { + if strings.TrimSpace(object) == "" { + return "batch" + } + return object +} + +func cleanGigaChatBatchFileID(fileID *string) *string { + if fileID == nil { + return nil + } + trimmed := strings.TrimSpace(*fileID) + if trimmed == "" { + return nil + } + return &trimmed +} + +func toBifrostGigaChatBatchOutputFileID(batch GigaChatBatch) *string { + if outputFileID := cleanGigaChatBatchFileID(batch.OutputFileID); outputFileID != nil { + return outputFileID + } + return cleanGigaChatBatchFileID(batch.ResultFileID) +} + +func (provider *GigaChatProvider) readGigaChatBatchOutputFile(ctx *schemas.BifrostContext, keys []schemas.Key, request *schemas.BifrostBatchResultsRequest, fileID string) (*schemas.BifrostFileContentResponse, *schemas.BifrostError) { + if len(keys) == 0 { + return nil, providerUtils.NewBifrostOperationError("no keys available to download GigaChat batch output file", nil) + } + + fileRequest := &schemas.BifrostFileContentRequest{ + Provider: provider.GetProviderKey(), + Model: request.Model, + FileID: strings.TrimSpace(fileID), + } + var lastErr *schemas.BifrostError + for _, key := range keys { + response, bifrostErr := provider.fileContentWithRefresh(ctx, key, fileRequest, false) + if isGigaChatUnauthorizedError(bifrostErr) { + response, bifrostErr = provider.fileContentWithRefresh(ctx, key, fileRequest, true) + } + if bifrostErr == nil { + return response, nil + } + lastErr = bifrostErr + } + return nil, lastErr +} + +func parseGigaChatBatchResultsJSONL(content []byte, logger schemas.Logger) ([]schemas.BatchResultItem, []schemas.BatchError) { + results := make([]schemas.BatchResultItem, 0) + parseResult := providerUtils.ParseJSONL(content, func(line []byte) error { + var resultRow GigaChatBatchResultRow + if err := json.Unmarshal(line, &resultRow); err != nil { + if logger != nil { + logger.Warn("failed to parse GigaChat batch result line: %v", err) + } + return err + } + + result := schemas.BatchResultItem{ + CustomID: strings.TrimSpace(resultRow.CustomID), + Response: resultRow.Response, + Result: resultRow.Result, + Error: resultRow.Error, + } + if result.CustomID == "" { + result.CustomID = strings.TrimSpace(resultRow.ID) + } + results = append(results, result) + return nil + }) + return results, parseResult.Errors +} + +func withGigaChatQuery(path string, values url.Values) string { + encoded := values.Encode() + if encoded == "" { + return path + } + if strings.Contains(path, "?") { + return path + "&" + encoded + } + return path + "?" + encoded +} + +func paginateGigaChatBatchList(response *schemas.BifrostBatchListResponse, nativeCursor string, limit int) (string, bool, error) { + offset := 0 + if nativeCursor != "" { + parsed, err := strconv.Atoi(nativeCursor) + if err != nil || parsed < 0 { + return "", false, fmt.Errorf("invalid GigaChat batch cursor %q", nativeCursor) + } + offset = parsed + } + if offset > len(response.Data) { + return "", false, fmt.Errorf("GigaChat batch cursor offset %d exceeds result count %d", offset, len(response.Data)) + } + + end := len(response.Data) + hasMore := false + if limit > 0 && limit < len(response.Data)-offset { + end = offset + limit + hasMore = true + } + response.Data = response.Data[offset:end] + response.FirstID = nil + response.LastID = nil + if len(response.Data) > 0 { + firstID := response.Data[0].ID + lastID := response.Data[len(response.Data)-1].ID + response.FirstID = &firstID + response.LastID = &lastID + } + if !hasMore { + return "", false, nil + } + return strconv.Itoa(end), true, nil +} + +func convertGigaChatBatchRequestItemsToJSONL(endpoint schemas.BatchEndpoint, requests []schemas.BatchRequestItem) ([]byte, error) { + var buf bytes.Buffer + for index, request := range requests { + row, err := toGigaChatBatchInputRowFromRequestItem(endpoint, request) + if err != nil { + return nil, fmt.Errorf("requests[%d]: %w", index, err) + } + if err := writeGigaChatBatchInputRow(&buf, row); err != nil { + return nil, fmt.Errorf("requests[%d]: %w", index, err) + } + } + return buf.Bytes(), nil +} + +func convertGigaChatBatchInputJSONL(endpoint schemas.BatchEndpoint, input []byte) ([]byte, error) { + var buf bytes.Buffer + parseResult := providerUtils.ParseJSONL(input, func(line []byte) error { + row, err := toGigaChatBatchInputRowFromJSONLine(endpoint, line) + if err != nil { + return err + } + return writeGigaChatBatchInputRow(&buf, row) + }) + if len(parseResult.Errors) > 0 { + return nil, formatGigaChatBatchJSONLErrors(parseResult.Errors) + } + return buf.Bytes(), nil +} + +func toGigaChatBatchInputRowFromRequestItem(defaultEndpoint schemas.BatchEndpoint, request schemas.BatchRequestItem) (GigaChatBatchInputRow, error) { + if len(request.Params) > 0 { + return GigaChatBatchInputRow{}, fmt.Errorf("params are not supported by GigaChat batch row conversion") + } + if request.Body == nil { + return GigaChatBatchInputRow{}, fmt.Errorf("body is required") + } + body, err := schemas.MarshalSorted(request.Body) + if err != nil { + return GigaChatBatchInputRow{}, fmt.Errorf("marshal body: %w", err) + } + row := openAICompatibleBatchInputRow{ + CustomID: request.CustomID, + Method: request.Method, + URL: request.URL, + Body: body, + } + return toGigaChatBatchInputRow(defaultEndpoint, row) +} + +func toGigaChatBatchInputRowFromJSONLine(defaultEndpoint schemas.BatchEndpoint, line []byte) (GigaChatBatchInputRow, error) { + var row openAICompatibleBatchInputRow + if err := json.Unmarshal(line, &row); err != nil { + return GigaChatBatchInputRow{}, fmt.Errorf("decode batch row: %w", err) + } + return toGigaChatBatchInputRow(defaultEndpoint, row) +} + +func toGigaChatBatchInputRow(defaultEndpoint schemas.BatchEndpoint, row openAICompatibleBatchInputRow) (GigaChatBatchInputRow, error) { + customID := strings.TrimSpace(row.CustomID) + if customID == "" { + return GigaChatBatchInputRow{}, fmt.Errorf("custom_id is required") + } + if method := strings.TrimSpace(row.Method); method != "" && !strings.EqualFold(method, http.MethodPost) { + return GigaChatBatchInputRow{}, fmt.Errorf("method %q is not supported by GigaChat batches", row.Method) + } + if len(bytes.TrimSpace(row.Body)) == 0 { + return GigaChatBatchInputRow{}, fmt.Errorf("body is required") + } + + endpoint := defaultEndpoint + if strings.TrimSpace(row.URL) != "" { + endpoint = schemas.BatchEndpoint(row.URL) + } + request, err := toGigaChatBatchRequestBody(endpoint, row.Body) + if err != nil { + return GigaChatBatchInputRow{}, err + } + return GigaChatBatchInputRow{ + ID: customID, + Request: request, + }, nil +} + +func toGigaChatBatchRequestBody(endpoint schemas.BatchEndpoint, body json.RawMessage) (json.RawMessage, error) { + switch normalizeGigaChatBatchEndpoint(string(endpoint)) { + case "/v1/chat/completions", "/chat/completions": + var request openaiProvider.OpenAIChatRequest + if err := json.Unmarshal(body, &request); err != nil { + return nil, fmt.Errorf("decode chat completion body: %w", err) + } + bifrostReq := request.ToBifrostChatRequest(gigaChatBatchConversionContext()) + if request.MaxTokens != nil { + if bifrostReq.Params == nil { + bifrostReq.Params = &schemas.ChatParameters{} + } + if bifrostReq.Params.MaxCompletionTokens == nil { + bifrostReq.Params.MaxCompletionTokens = request.MaxTokens + } + } + bifrostReq.Provider = schemas.GigaChat + gigaChatReq, err := ToGigaChatChatRequest(gigaChatBatchConversionContext(), bifrostReq) + if err != nil { + return nil, err + } + return marshalGigaChatBatchRequest(gigaChatReq) + case "/v1/responses", "/responses": + var request openaiProvider.OpenAIResponsesRequest + if err := json.Unmarshal(body, &request); err != nil { + return nil, fmt.Errorf("decode responses body: %w", err) + } + bifrostReq := request.ToBifrostResponsesRequest(gigaChatBatchConversionContext()) + bifrostReq.Provider = schemas.GigaChat + gigaChatReq, err := ToGigaChatResponsesRequest(bifrostReq) + if err != nil { + return nil, err + } + return marshalGigaChatBatchRequest(gigaChatReq) + case "/v1/embeddings", "/embeddings": + var request openaiProvider.OpenAIEmbeddingRequest + if err := json.Unmarshal(body, &request); err != nil { + return nil, fmt.Errorf("decode embeddings body: %w", err) + } + bifrostReq := request.ToBifrostEmbeddingRequest(gigaChatBatchConversionContext()) + bifrostReq.Provider = schemas.GigaChat + gigaChatReq, err := ToGigaChatEmbeddingRequest(bifrostReq) + if err != nil { + return nil, err + } + return marshalGigaChatBatchRequest(gigaChatReq) + default: + return nil, fmt.Errorf("GigaChat batches do not support endpoint %q", endpoint) + } +} + +func marshalGigaChatBatchRequest(request providerUtils.RequestBodyWithExtraParams) (json.RawMessage, error) { + body, err := schemas.MarshalSorted(request) + if err != nil { + return nil, fmt.Errorf("marshal GigaChat batch request: %w", err) + } + if extraParams := request.GetExtraParams(); len(extraParams) > 0 { + body, err = providerUtils.MergeExtraParamsIntoJSON(body, extraParams) + if err != nil { + return nil, fmt.Errorf("merge GigaChat batch request extra params: %w", err) + } + } + var compacted bytes.Buffer + if err := json.Compact(&compacted, body); err != nil { + return nil, fmt.Errorf("compact GigaChat batch request: %w", err) + } + return json.RawMessage(compacted.Bytes()), nil +} + +func writeGigaChatBatchInputRow(buf *bytes.Buffer, row GigaChatBatchInputRow) error { + line, err := schemas.MarshalSorted(row) + if err != nil { + return err + } + buf.Write(line) + buf.WriteByte('\n') + return nil +} + +func formatGigaChatBatchJSONLErrors(errors []schemas.BatchError) error { + if len(errors) == 0 { + return nil + } + messages := make([]string, 0, len(errors)) + for _, parseErr := range errors { + if parseErr.Line != nil { + messages = append(messages, fmt.Sprintf("line %d: %s", *parseErr.Line, parseErr.Message)) + } else { + messages = append(messages, parseErr.Message) + } + } + return fmt.Errorf("failed to convert GigaChat batch JSONL: %s", strings.Join(messages, "; ")) +} + +func gigaChatBatchConversionContext() *schemas.BifrostContext { + return schemas.NewBifrostContext(nil, schemas.NoDeadline) +} diff --git a/core/providers/gigachat/chat.go b/core/providers/gigachat/chat.go new file mode 100644 index 00000000000..5950586c3f7 --- /dev/null +++ b/core/providers/gigachat/chat.go @@ -0,0 +1,837 @@ +package gigachat + +import ( + "bytes" + "encoding/json" + "fmt" + "net/http" + "sort" + "strings" + + providerUtils "github.com/maximhq/bifrost/core/providers/utils" + schemas "github.com/maximhq/bifrost/core/schemas" +) + +const ( + gigaChatMinReasoningMaxTokens = 1 + gigaChatDefaultCompletionMaxTokens = 4096 +) + +// ToGigaChatChatRequest converts a Bifrost chat request to GigaChat v1 format. +func ToGigaChatChatRequest(_ *schemas.BifrostContext, bifrostReq *schemas.BifrostChatRequest) (*GigaChatChatRequest, error) { + if bifrostReq == nil { + return nil, fmt.Errorf("bifrost chat request is nil") + } + if strings.TrimSpace(bifrostReq.Model) == "" { + return nil, fmt.Errorf("model is required") + } + if len(bifrostReq.Input) == 0 { + return nil, fmt.Errorf("messages are required") + } + + toolCallNamesByID := collectGigaChatChatToolCallNames(bifrostReq.Input) + messages := make([]GigaChatChatMessage, 0, len(bifrostReq.Input)) + needsAutoFunctionCall := false + for index, message := range bifrostReq.Input { + convertedMessage, messageNeedsAutoFunctionCall, err := toGigaChatChatMessage(message, toolCallNamesByID) + if err != nil { + return nil, fmt.Errorf("messages[%d]: %w", index, err) + } + messages = append(messages, convertedMessage) + needsAutoFunctionCall = needsAutoFunctionCall || messageNeedsAutoFunctionCall + } + + gigaChatReq := &GigaChatChatRequest{ + Model: bifrostReq.Model, + Messages: messages, + Stream: schemas.Ptr(false), + } + if bifrostReq.Params == nil { + if needsAutoFunctionCall { + gigaChatReq.FunctionCall = "auto" + } + return gigaChatReq, nil + } + + if unsupportedParams := unsupportedGigaChatChatParams(bifrostReq.Params); len(unsupportedParams) > 0 { + return nil, fmt.Errorf("GigaChat v1 chat completions do not support parameter(s): %s", strings.Join(unsupportedParams, ", ")) + } + + gigaChatReq.Temperature = bifrostReq.Params.Temperature + gigaChatReq.TopP = bifrostReq.Params.TopP + gigaChatReq.MaxTokens = bifrostReq.Params.MaxCompletionTokens + gigaChatReq.N = bifrostReq.Params.N + gigaChatReq.Stop = bifrostReq.Params.Stop + gigaChatReq.ReasoningEffort = toGigaChatChatReasoningEffort(bifrostReq.Model, bifrostReq.Params) + gigaChatReq.ExtraParams = bifrostReq.Params.ExtraParams + responseFormat, err := toGigaChatChatResponseFormat(bifrostReq.Params.ResponseFormat) + if err != nil { + return nil, err + } + gigaChatReq.ResponseFormat = responseFormat + functions, functionNames, err := toGigaChatChatFunctions(bifrostReq.Params.Tools) + if err != nil { + return nil, err + } + gigaChatReq.Functions = functions + functionCall, err := toGigaChatChatFunctionCall(bifrostReq.Params.ToolChoice, functionNames) + if err != nil { + return nil, err + } + gigaChatReq.FunctionCall = functionCall + if needsAutoFunctionCall && gigaChatReq.FunctionCall == nil { + gigaChatReq.FunctionCall = "auto" + } + + return gigaChatReq, nil +} + +// ToGigaChatChatStreamRequest converts a Bifrost chat request to a streaming GigaChat v1 request. +func ToGigaChatChatStreamRequest(ctx *schemas.BifrostContext, bifrostReq *schemas.BifrostChatRequest) (*GigaChatChatRequest, error) { + gigaChatReq, err := ToGigaChatChatRequest(ctx, bifrostReq) + if err != nil { + return nil, err + } + gigaChatReq.Stream = schemas.Ptr(true) + return gigaChatReq, nil +} + +// ToBifrostChatResponse converts a GigaChat v1 chat response to Bifrost format. +func ToBifrostChatResponse(providerName schemas.ModelProvider, response *GigaChatChatResponse) *schemas.BifrostChatResponse { + if response == nil { + return nil + } + + choices := make([]schemas.BifrostResponseChoice, 0, len(response.Choices)) + for _, choice := range response.Choices { + choices = append(choices, schemas.BifrostResponseChoice{ + Index: choice.Index, + FinishReason: toBifrostGigaChatFinishReason(choice.FinishReason), + LogProbs: choice.LogProbs, + ChatNonStreamResponseChoice: &schemas.ChatNonStreamResponseChoice{Message: toBifrostGigaChatMessage(choice.Message)}, + }) + } + + return &schemas.BifrostChatResponse{ + ID: response.ID, + Choices: choices, + Created: response.Created, + Model: response.Model, + Object: response.Object, + SystemFingerprint: response.SystemFingerprint, + Usage: toBifrostGigaChatUsage(response.Usage), + ExtraParams: response.ExtraParams, + ExtraFields: schemas.BifrostResponseExtraFields{ + Provider: providerName, + }, + } +} + +// ToBifrostChatStreamResponse converts a GigaChat v1 chat SSE chunk to Bifrost format. +func ToBifrostChatStreamResponse(providerName schemas.ModelProvider, response *GigaChatChatStreamResponse) *schemas.BifrostChatResponse { + if response == nil { + return nil + } + + choices := make([]schemas.BifrostResponseChoice, 0, len(response.Choices)) + for _, choice := range response.Choices { + choices = append(choices, schemas.BifrostResponseChoice{ + Index: choice.Index, + FinishReason: toBifrostGigaChatFinishReason(choice.FinishReason), + LogProbs: choice.LogProbs, + ChatStreamResponseChoice: &schemas.ChatStreamResponseChoice{ + Delta: toBifrostGigaChatStreamDelta(choice.Index, choice.Delta), + }, + }) + } + + return &schemas.BifrostChatResponse{ + ID: response.ID, + Choices: choices, + Created: response.Created, + Model: response.Model, + Object: toBifrostGigaChatChatStreamObject(response.Object), + SystemFingerprint: response.SystemFingerprint, + Usage: toBifrostGigaChatUsage(response.Usage), + ExtraParams: response.ExtraParams, + ExtraFields: schemas.BifrostResponseExtraFields{ + Provider: providerName, + }, + } +} + +func toBifrostGigaChatChatStreamObject(object string) string { + switch strings.TrimSpace(object) { + case "", "chat.completion", "chat.completions": + return "chat.completion.chunk" + default: + return object + } +} + +func handleGigaChatChatStreamResponse(providerName schemas.ModelProvider) func([]byte, *schemas.BifrostChatResponse, []byte, bool, bool) (interface{}, interface{}, *schemas.BifrostError) { + return func(responseBody []byte, response *schemas.BifrostChatResponse, requestBody []byte, sendBackRawRequest bool, sendBackRawResponse bool) (interface{}, interface{}, *schemas.BifrostError) { + if bifrostErr := parseGigaChatStreamError(responseBody, providerName); bifrostErr != nil { + rawRequest, rawResponse, _ := providerUtils.HandleProviderResponse(responseBody, &GigaChatErrorResponse{}, requestBody, sendBackRawRequest, sendBackRawResponse) + return redactGigaChatRawValue(rawRequest), redactGigaChatRawValue(rawResponse), bifrostErr + } + + var gigaChatResponse GigaChatChatStreamResponse + rawRequest, rawResponse, bifrostErr := providerUtils.HandleProviderResponse(responseBody, &gigaChatResponse, requestBody, sendBackRawRequest, sendBackRawResponse) + if bifrostErr != nil { + return rawRequest, rawResponse, bifrostErr + } + + converted := ToBifrostChatStreamResponse(providerName, &gigaChatResponse) + if converted == nil { + return rawRequest, rawResponse, newGigaChatProviderResponseError("GigaChat chat completion stream response is empty", nil) + } + *response = *converted + return rawRequest, rawResponse, nil + } +} + +func withGigaChatChatResponseProvider(providerName schemas.ModelProvider) func(*schemas.BifrostChatResponse) *schemas.BifrostChatResponse { + return func(response *schemas.BifrostChatResponse) *schemas.BifrostChatResponse { + if response != nil { + response.ExtraFields.Provider = providerName + } + return response + } +} + +func toGigaChatChatMessage(message schemas.ChatMessage, toolCallNamesByID map[string]string) (GigaChatChatMessage, bool, error) { + switch message.Role { + case schemas.ChatMessageRoleSystem, schemas.ChatMessageRoleUser, schemas.ChatMessageRoleAssistant: + case schemas.ChatMessageRoleTool: + convertedMessage, err := toGigaChatFunctionResultMessage(message, toolCallNamesByID) + return convertedMessage, false, err + case schemas.ChatMessageRoleDeveloper: + return GigaChatChatMessage{}, false, fmt.Errorf("developer messages are not supported by GigaChat v1 chat completions") + default: + return GigaChatChatMessage{}, false, fmt.Errorf("unsupported role %q", message.Role) + } + if message.ChatToolMessage != nil { + return GigaChatChatMessage{}, false, fmt.Errorf("tool message fields are not supported by GigaChat v1 chat completions") + } + if message.ChatAssistantMessage != nil { + if len(message.ChatAssistantMessage.ToolCalls) > 0 { + convertedMessage, err := toGigaChatAssistantFunctionCallMessage(message) + return convertedMessage, false, err + } + if message.ChatAssistantMessage.Refusal != nil || + message.ChatAssistantMessage.Audio != nil || + len(message.ChatAssistantMessage.Annotations) > 0 { + return GigaChatChatMessage{}, false, fmt.Errorf("assistant-only OpenAI metadata is not supported by GigaChat v1 chat completions") + } + } + reasoning, err := toGigaChatChatReasoningContent(message.ChatAssistantMessage) + if err != nil { + return GigaChatChatMessage{}, false, err + } + + content, attachments, needsAutoFunctionCall, err := toGigaChatChatMessageContent(message.Content) + if err != nil { + return GigaChatChatMessage{}, false, err + } + if content == nil && len(attachments) > 0 { + content = &schemas.ChatMessageContent{ContentStr: schemas.Ptr("")} + } + + return GigaChatChatMessage{ + Role: string(message.Role), + Content: content, + Attachments: attachments, + Name: message.Name, + Reasoning: reasoning, + }, needsAutoFunctionCall, nil +} + +func collectGigaChatChatToolCallNames(messages []schemas.ChatMessage) map[string]string { + toolCallNamesByID := make(map[string]string) + for _, message := range messages { + if message.ChatAssistantMessage == nil { + continue + } + for _, toolCall := range message.ChatAssistantMessage.ToolCalls { + if toolCall.ID == nil || strings.TrimSpace(*toolCall.ID) == "" || toolCall.Function.Name == nil || strings.TrimSpace(*toolCall.Function.Name) == "" { + continue + } + toolCallNamesByID[strings.TrimSpace(*toolCall.ID)] = strings.TrimSpace(*toolCall.Function.Name) + } + } + return toolCallNamesByID +} + +func toGigaChatAssistantFunctionCallMessage(message schemas.ChatMessage) (GigaChatChatMessage, error) { + if message.ChatAssistantMessage == nil || len(message.ChatAssistantMessage.ToolCalls) == 0 { + return GigaChatChatMessage{}, fmt.Errorf("assistant function_call is required") + } + if len(message.ChatAssistantMessage.ToolCalls) > 1 { + return GigaChatChatMessage{}, fmt.Errorf("GigaChat v1 chat completions support one function call per assistant message") + } + toolCall := message.ChatAssistantMessage.ToolCalls[0] + if toolCall.Type != nil && *toolCall.Type != "" && *toolCall.Type != string(schemas.ChatToolTypeFunction) { + return GigaChatChatMessage{}, fmt.Errorf("assistant tool call type %q is not supported by GigaChat v1 chat completions", *toolCall.Type) + } + if toolCall.Function.Name == nil || strings.TrimSpace(*toolCall.Function.Name) == "" { + return GigaChatChatMessage{}, fmt.Errorf("assistant function_call name is required") + } + arguments, err := parseGigaChatChatFunctionArguments(toolCall.Function.Arguments) + if err != nil { + return GigaChatChatMessage{}, err + } + + content, attachments, _, err := toGigaChatChatMessageContent(message.Content) + if err != nil { + return GigaChatChatMessage{}, err + } + if len(attachments) > 0 { + return GigaChatChatMessage{}, fmt.Errorf("assistant function_call messages do not support attachments") + } + if content == nil { + content = &schemas.ChatMessageContent{ContentStr: schemas.Ptr("")} + } + reasoning, err := toGigaChatChatReasoningContent(message.ChatAssistantMessage) + if err != nil { + return GigaChatChatMessage{}, err + } + + return GigaChatChatMessage{ + Role: string(schemas.ChatMessageRoleAssistant), + Content: content, + Name: message.Name, + Reasoning: reasoning, + FunctionCall: &GigaChatFunctionCall{ + Name: strings.TrimSpace(*toolCall.Function.Name), + Arguments: arguments, + }, + FunctionsStateID: toolCall.ID, + }, nil +} + +func toGigaChatFunctionResultMessage(message schemas.ChatMessage, toolCallNamesByID map[string]string) (GigaChatChatMessage, error) { + if message.ChatToolMessage == nil { + return GigaChatChatMessage{}, fmt.Errorf("function result message requires tool message fields") + } + name := "" + if message.Name != nil { + name = strings.TrimSpace(*message.Name) + } + if name == "" && message.ChatToolMessage.ToolCallID != nil { + name = toolCallNamesByID[strings.TrimSpace(*message.ChatToolMessage.ToolCallID)] + } + if name == "" { + return GigaChatChatMessage{}, fmt.Errorf("function result message requires function name or matching tool_call_id") + } + content, attachments, _, err := toGigaChatChatMessageContent(message.Content) + if err != nil { + return GigaChatChatMessage{}, err + } + if len(attachments) > 0 { + return GigaChatChatMessage{}, fmt.Errorf("function result messages do not support attachments") + } + if content == nil || content.ContentStr == nil || strings.TrimSpace(*content.ContentStr) == "" { + return GigaChatChatMessage{}, fmt.Errorf("function result message content is required") + } + trimmedContent := bytes.TrimSpace([]byte(*content.ContentStr)) + if !json.Valid(trimmedContent) || len(trimmedContent) == 0 || trimmedContent[0] != '{' { + return GigaChatChatMessage{}, fmt.Errorf("function result message content must be a JSON object string") + } + + return GigaChatChatMessage{ + Role: "function", + Content: content, + Name: &name, + }, nil +} + +func toGigaChatChatReasoningContent(assistantMessage *schemas.ChatAssistantMessage) (*string, error) { + if assistantMessage == nil { + return nil, nil + } + if assistantMessage.Reasoning != nil { + return assistantMessage.Reasoning, nil + } + if len(assistantMessage.ReasoningDetails) == 0 { + return nil, nil + } + + var reasoningBuilder strings.Builder + for _, detail := range assistantMessage.ReasoningDetails { + var text *string + switch detail.Type { + case schemas.BifrostReasoningDetailsTypeText: + text = detail.Text + case schemas.BifrostReasoningDetailsTypeSummary: + text = detail.Summary + default: + return nil, fmt.Errorf("assistant reasoning detail type %q is not supported by GigaChat v1 chat completions", detail.Type) + } + if text == nil { + return nil, fmt.Errorf("assistant reasoning detail type %q requires text content for GigaChat v1 chat completions", detail.Type) + } + reasoningBuilder.WriteString(*text) + } + reasoning := reasoningBuilder.String() + return &reasoning, nil +} + +func toGigaChatChatMessageContent(content *schemas.ChatMessageContent) (*schemas.ChatMessageContent, []string, bool, error) { + if content == nil { + return nil, nil, false, nil + } + if content.ContentStr != nil { + return content, nil, false, nil + } + if len(content.ContentBlocks) == 0 { + return content, nil, false, nil + } + + var textBuilder strings.Builder + attachments := make([]string, 0) + needsAutoFunctionCall := false + for index, block := range content.ContentBlocks { + switch block.Type { + case schemas.ChatContentBlockTypeText: + if block.Text != nil { + textBuilder.WriteString(*block.Text) + } + case schemas.ChatContentBlockTypeFile: + attachmentID, blockNeedsAutoFunctionCall, err := toGigaChatChatAttachment(index, block) + if err != nil { + return nil, nil, false, err + } + attachments = append(attachments, attachmentID) + needsAutoFunctionCall = needsAutoFunctionCall || blockNeedsAutoFunctionCall + case schemas.ChatContentBlockTypeImage: + return nil, nil, false, fmt.Errorf("content block %d: image_url must be uploaded before GigaChat v1 chat completions request conversion", index) + default: + return nil, nil, false, fmt.Errorf("content block %d with type %q is not supported by GigaChat v1 chat completions", index, block.Type) + } + } + text := textBuilder.String() + return &schemas.ChatMessageContent{ContentStr: &text}, attachments, needsAutoFunctionCall, nil +} + +func toGigaChatChatAttachment(index int, block schemas.ChatContentBlock) (string, bool, error) { + if block.File == nil { + return "", false, fmt.Errorf("content block %d: file block is missing file payload", index) + } + if block.File.FileData != nil || block.File.FileURL != nil { + return "", false, fmt.Errorf("content block %d: GigaChat v1 chat completions supports pre-uploaded file_id references only; upload inline file content before request conversion", index) + } + if block.File.FileID == nil || strings.TrimSpace(*block.File.FileID) == "" { + return "", false, fmt.Errorf("content block %d: GigaChat attachment requires file_id", index) + } + return strings.TrimSpace(*block.File.FileID), gigaChatChatFileRequiresAutoFunctionCall(block.File), nil +} + +func unsupportedGigaChatChatParams(params *schemas.ChatParameters) []string { + if params == nil { + return nil + } + + unsupported := make([]string, 0) + addIf := func(condition bool, name string) { + if condition { + unsupported = append(unsupported, name) + } + } + + addIf(params.Audio != nil, "audio") + addIf(params.FrequencyPenalty != nil, "frequency_penalty") + addIf(params.LogitBias != nil, "logit_bias") + addIf(params.LogProbs != nil && *params.LogProbs, "logprobs") + addIf(params.Metadata != nil && len(*params.Metadata) > 0, "metadata") + addIf(len(params.Modalities) > 0, "modalities") + addIf(params.ParallelToolCalls != nil && *params.ParallelToolCalls, "parallel_tool_calls") + addIf(params.Prediction != nil, "prediction") + addIf(params.PresencePenalty != nil, "presence_penalty") + addIf(params.PromptCacheKey != nil, "prompt_cache_key") + addIf(params.PromptCacheRetention != nil, "prompt_cache_retention") + addIf(params.SafetyIdentifier != nil, "safety_identifier") + addIf(params.Seed != nil, "seed") + addIf(params.ServiceTier != nil, "service_tier") + addIf(params.StreamOptions != nil, "stream_options") + addIf(params.Store != nil && *params.Store, "store") + addIf(params.TopLogProbs != nil, "top_logprobs") + addIf(params.User != nil, "user") + addIf(params.Verbosity != nil, "verbosity") + addIf(params.WebSearchOptions != nil, "web_search_options") + addIf(params.TopK != nil, "top_k") + addIf(params.Speed != nil, "speed") + addIf(params.InferenceGeo != nil, "inference_geo") + addIf(len(params.MCPServers) > 0, "mcp_servers") + addIf(params.Container != nil, "container") + addIf(params.CacheControl != nil, "cache_control") + addIf(params.TaskBudget != nil, "task_budget") + addIf(len(bytes.TrimSpace(params.ContextManagement)) > 0, "context_management") + unsupported = append(unsupported, unsupportedGigaChatToolControlExtraParams(params.ExtraParams, "functions", "function_call", "tools", "tool_config", "parallel_tool_calls")...) + + sort.Strings(unsupported) + return unsupported +} + +func toGigaChatChatResponseFormat(responseFormat *interface{}) (interface{}, error) { + if responseFormat == nil { + return nil, nil + } + + responseFormatMap, ok := schemas.SafeExtractOrderedMap(*responseFormat) + if !ok || responseFormatMap == nil { + return nil, fmt.Errorf("response_format must be a JSON object") + } + + formatTypeRaw, ok := responseFormatMap.Get("type") + if !ok { + return nil, fmt.Errorf("response_format.type is required") + } + formatType, ok := schemas.SafeExtractString(formatTypeRaw) + if !ok || strings.TrimSpace(formatType) == "" { + return nil, fmt.Errorf("response_format.type must be a non-empty string") + } + formatType = strings.TrimSpace(formatType) + + switch formatType { + case "json_schema": + return toGigaChatChatJSONSchemaResponseFormat(responseFormatMap) + default: + return nil, fmt.Errorf("response_format type %q is not supported by GigaChat v1 chat completions", formatType) + } +} + +func toGigaChatChatJSONSchemaResponseFormat(responseFormatMap *schemas.OrderedMap) (interface{}, error) { + var ( + schemaRaw interface{} + name *string + description *string + strict *bool + err error + ) + + if jsonSchemaRaw, ok := responseFormatMap.Get("json_schema"); ok { + jsonSchemaMap, ok := schemas.SafeExtractOrderedMap(jsonSchemaRaw) + if !ok || jsonSchemaMap == nil { + return nil, fmt.Errorf("response_format json_schema must be a JSON object") + } + schemaRaw, ok = jsonSchemaMap.Get("schema") + if !ok || schemaRaw == nil { + return nil, fmt.Errorf("response_format json_schema requires schema") + } + name, err = optionalGigaChatResponseFormatString(jsonSchemaMap, "name") + if err != nil { + return nil, err + } + description, err = optionalGigaChatResponseFormatString(jsonSchemaMap, "description") + if err != nil { + return nil, err + } + strict, err = optionalGigaChatResponseFormatBool(jsonSchemaMap, "strict") + if err != nil { + return nil, err + } + } else { + var ok bool + schemaRaw, ok = responseFormatMap.Get("schema") + if !ok || schemaRaw == nil { + return nil, fmt.Errorf("response_format json_schema requires schema") + } + name, err = optionalGigaChatResponseFormatString(responseFormatMap, "name") + if err != nil { + return nil, err + } + description, err = optionalGigaChatResponseFormatString(responseFormatMap, "description") + if err != nil { + return nil, err + } + strict, err = optionalGigaChatResponseFormatBool(responseFormatMap, "strict") + if err != nil { + return nil, err + } + } + + schemaMap, ok := asGigaChatSchemaMap(schemaRaw) + if !ok || schemaMap == nil { + return nil, fmt.Errorf("response_format json_schema.schema must be a JSON object") + } + schema, err := cloneGigaChatSchemaMap(schemaMap) + if err != nil { + return nil, fmt.Errorf("response_format json_schema.schema is invalid: %w", err) + } + schemaWithMetadata := withGigaChatResponseFormatSchemaMetadata(schema, name, description) + + gigaChatResponseFormat := schemas.NewOrderedMapFromPairs( + schemas.KV("type", "json_schema"), + schemas.KV("schema", schemaWithMetadata), + ) + if strict != nil { + gigaChatResponseFormat.Set("strict", *strict) + } + return gigaChatResponseFormat, nil +} + +func optionalGigaChatResponseFormatString(values *schemas.OrderedMap, name string) (*string, error) { + raw, ok := values.Get(name) + if !ok || raw == nil { + return nil, nil + } + value, ok := schemas.SafeExtractString(raw) + if !ok { + return nil, fmt.Errorf("response_format json_schema.%s must be a string", name) + } + value = strings.TrimSpace(value) + if value == "" { + return nil, nil + } + return &value, nil +} + +func optionalGigaChatResponseFormatBool(values *schemas.OrderedMap, name string) (*bool, error) { + raw, ok := values.Get(name) + if !ok || raw == nil { + return nil, nil + } + value, ok := schemas.SafeExtractBool(raw) + if !ok { + return nil, fmt.Errorf("response_format json_schema.%s must be a boolean", name) + } + return &value, nil +} + +func toGigaChatChatReasoningEffort(model string, params *schemas.ChatParameters) *string { + if params == nil || params.Reasoning == nil { + return nil + } + if params.Reasoning.Enabled != nil && !*params.Reasoning.Enabled { + return nil + } + if params.Reasoning.Effort != nil { + effort := normalizeGigaChatChatReasoningEffort(*params.Reasoning.Effort) + if effort == "" || effort == "none" { + return nil + } + return &effort + } + if params.Reasoning.MaxTokens != nil { + maxCompletionTokens := providerUtils.GetMaxOutputTokensOrDefault(schemas.GigaChat, model, gigaChatDefaultCompletionMaxTokens) + if params.MaxCompletionTokens != nil { + maxCompletionTokens = *params.MaxCompletionTokens + } + effort := providerUtils.GetReasoningEffortFromBudgetTokens(*params.Reasoning.MaxTokens, gigaChatMinReasoningMaxTokens, maxCompletionTokens) + if effort == "none" { + return nil + } + return &effort + } + return nil +} + +func normalizeGigaChatChatReasoningEffort(effort string) string { + normalized := strings.TrimSpace(strings.ToLower(effort)) + switch normalized { + case "minimal": + return "low" + case "xhigh", "max": + return "high" + default: + return normalized + } +} + +func parseGigaChatChatFunctionArguments(arguments string) (json.RawMessage, error) { + trimmed := bytes.TrimSpace([]byte(arguments)) + if len(trimmed) == 0 { + return json.RawMessage(`{}`), nil + } + if !json.Valid(trimmed) || trimmed[0] != '{' { + return nil, fmt.Errorf("function_call arguments must be a JSON object") + } + var compacted bytes.Buffer + if err := json.Compact(&compacted, trimmed); err != nil { + return nil, fmt.Errorf("function_call arguments must be valid JSON: %w", err) + } + return json.RawMessage(compacted.Bytes()), nil +} + +func toBifrostGigaChatMessage(message *GigaChatChatMessage) *schemas.ChatMessage { + if message == nil { + return nil + } + + role := schemas.ChatMessageRole(message.Role) + if role == "" { + role = schemas.ChatMessageRoleAssistant + } + + bifrostMessage := &schemas.ChatMessage{ + Role: role, + Content: message.Content, + Name: message.Name, + } + var assistantMessage *schemas.ChatAssistantMessage + if message.Reasoning != nil { + assistantMessage = &schemas.ChatAssistantMessage{ + Reasoning: message.Reasoning, + ReasoningDetails: toBifrostGigaChatReasoningDetails(message.Reasoning), + } + } + if message.FunctionCall != nil { + arguments := compactGigaChatFunctionArguments(message.FunctionCall.Arguments) + toolCallType := string(schemas.ChatToolTypeFunction) + toolCall := schemas.ChatAssistantMessageToolCall{ + Type: &toolCallType, + ID: message.FunctionsStateID, + Function: schemas.ChatAssistantMessageToolCallFunction{ + Name: &message.FunctionCall.Name, + Arguments: arguments, + }, + } + if assistantMessage == nil { + assistantMessage = &schemas.ChatAssistantMessage{} + } + assistantMessage.ToolCalls = []schemas.ChatAssistantMessageToolCall{toolCall} + } + if assistantMessage != nil { + bifrostMessage.ChatAssistantMessage = assistantMessage + } + return bifrostMessage +} + +func toBifrostGigaChatStreamDelta(_ int, delta *GigaChatChatStreamDelta) *schemas.ChatStreamResponseChoiceDelta { + if delta == nil { + return &schemas.ChatStreamResponseChoiceDelta{} + } + + bifrostDelta := &schemas.ChatStreamResponseChoiceDelta{ + Role: delta.Role, + Content: delta.Content, + Reasoning: delta.Reasoning, + ReasoningDetails: toBifrostGigaChatReasoningDetails(delta.Reasoning), + } + if delta.FunctionCall != nil { + arguments := compactGigaChatFunctionArguments(delta.FunctionCall.Arguments) + toolCallType := string(schemas.ChatToolTypeFunction) + bifrostDelta.ToolCalls = []schemas.ChatAssistantMessageToolCall{ + { + Index: 0, + Type: &toolCallType, + ID: delta.FunctionsStateID, + Function: schemas.ChatAssistantMessageToolCallFunction{ + Name: &delta.FunctionCall.Name, + Arguments: arguments, + }, + }, + } + } + return bifrostDelta +} + +func toBifrostGigaChatReasoningDetails(reasoning *string) []schemas.ChatReasoningDetails { + if reasoning == nil { + return nil + } + text := *reasoning + return []schemas.ChatReasoningDetails{ + { + Index: 0, + Type: schemas.BifrostReasoningDetailsTypeText, + Text: &text, + }, + } +} + +func compactGigaChatFunctionArguments(arguments json.RawMessage) string { + if len(arguments) == 0 { + return "" + } + var compacted bytes.Buffer + if err := json.Compact(&compacted, arguments); err != nil { + return string(arguments) + } + return compacted.String() +} + +func toBifrostGigaChatFinishReason(finishReason *string) *string { + if finishReason == nil { + return nil + } + if *finishReason == "function_call" { + return schemas.Ptr(string(schemas.BifrostFinishReasonToolCalls)) + } + return finishReason +} + +func toBifrostGigaChatUsage(usage *GigaChatChatUsage) *schemas.BifrostLLMUsage { + if usage == nil { + return nil + } + promptTokens := usage.PromptTokens + if promptTokens == 0 && usage.InputTokens > 0 { + promptTokens = usage.InputTokens + } + completionTokens := usage.CompletionTokens + if completionTokens == 0 && usage.OutputTokens > 0 { + completionTokens = usage.OutputTokens + } + totalTokens := usage.TotalTokens + if totalTokens == 0 && promptTokens+completionTokens > 0 { + totalTokens = promptTokens + completionTokens + } + + bifrostUsage := &schemas.BifrostLLMUsage{ + PromptTokens: promptTokens, + CompletionTokens: completionTokens, + TotalTokens: totalTokens, + } + cachedTokens := usage.PrecachedPromptTokens + if cachedTokens == 0 && usage.InputTokensDetails != nil { + cachedTokens = usage.InputTokensDetails.CachedTokens + if cachedTokens == 0 { + cachedTokens = usage.InputTokensDetails.CachedReadTokens + } + } + if cachedTokens > 0 { + bifrostUsage.PromptTokensDetails = &schemas.ChatPromptTokensDetails{ + CachedReadTokens: cachedTokens, + } + } + return bifrostUsage +} + +func parseGigaChatStreamError(responseBody []byte, providerName schemas.ModelProvider) *schemas.BifrostError { + var errorResp GigaChatErrorResponse + if err := json.Unmarshal(responseBody, &errorResp); err != nil { + return nil + } + if errorResp.Status == nil && errorResp.Code == nil && gigaChatErrorMessage(errorResp) == "" { + return nil + } + + statusCode := http.StatusBadGateway + if errorResp.Status != nil { + statusCode = *errorResp.Status + } + + bifrostErr := &schemas.BifrostError{ + IsBifrostError: false, + StatusCode: &statusCode, + Error: &schemas.ErrorField{}, + ExtraFields: schemas.BifrostErrorExtraFields{ + Provider: providerName, + }, + } + if message := gigaChatErrorMessage(errorResp); message != "" { + bifrostErr.Error.Message = redactGigaChatSensitiveText(message) + } else { + bifrostErr.Error.Message = fmt.Sprintf("GigaChat API error (status %d)", statusCode) + } + if codeValue, ok := gigaChatErrorCode(errorResp); ok { + code := codeValue + bifrostErr.Error.Code = &code + } else if errorResp.Status != nil { + code := fmt.Sprintf("%d", *errorResp.Status) + bifrostErr.Error.Code = &code + } + return bifrostErr +} diff --git a/core/providers/gigachat/chat_attachments.go b/core/providers/gigachat/chat_attachments.go new file mode 100644 index 00000000000..cb857885851 --- /dev/null +++ b/core/providers/gigachat/chat_attachments.go @@ -0,0 +1,377 @@ +package gigachat + +import ( + "encoding/base64" + "fmt" + "mime" + "net/url" + "path/filepath" + "strings" + + providerUtils "github.com/maximhq/bifrost/core/providers/utils" + "github.com/maximhq/bifrost/core/schemas" +) + +func (provider *GigaChatProvider) prepareGigaChatChatAttachments(ctx *schemas.BifrostContext, key schemas.Key, request *schemas.BifrostChatRequest) (*schemas.BifrostChatRequest, *schemas.BifrostError) { + if request == nil { + return nil, providerUtils.NewBifrostOperationError("chat completion request is nil", nil) + } + + var prepared *schemas.BifrostChatRequest + for messageIndex := range request.Input { + content := request.Input[messageIndex].Content + if content == nil || len(content.ContentBlocks) == 0 { + continue + } + + for blockIndex, block := range content.ContentBlocks { + if gigaChatChatAttachmentMayUpload(block) { + if replacement, ok := provider.getCachedGigaChatChatAttachment(ctx, key, request, messageIndex, blockIndex); ok { + if prepared == nil { + prepared = cloneGigaChatChatRequestForAttachmentUpload(request) + } + prepared.Input[messageIndex].Content.ContentBlocks[blockIndex] = replacement + continue + } + } + + replacement, changed, bifrostErr := provider.prepareGigaChatChatAttachmentBlock(ctx, key, blockIndex, block) + if bifrostErr != nil { + return nil, bifrostErr + } + if !changed { + continue + } + + if prepared == nil { + prepared = cloneGigaChatChatRequestForAttachmentUpload(request) + } + provider.setCachedGigaChatChatAttachment(ctx, key, request, messageIndex, blockIndex, replacement) + prepared.Input[messageIndex].Content.ContentBlocks[blockIndex] = replacement + } + } + + if prepared != nil { + return prepared, nil + } + return request, nil +} + +func gigaChatChatAttachmentMayUpload(block schemas.ChatContentBlock) bool { + switch block.Type { + case schemas.ChatContentBlockTypeImage: + return true + case schemas.ChatContentBlockTypeFile: + if block.File == nil { + return false + } + return block.File.FileID == nil || strings.TrimSpace(*block.File.FileID) == "" + default: + return false + } +} + +func cloneGigaChatChatRequestForAttachmentUpload(request *schemas.BifrostChatRequest) *schemas.BifrostChatRequest { + prepared := *request + prepared.Input = make([]schemas.ChatMessage, len(request.Input)) + copy(prepared.Input, request.Input) + + for i := range prepared.Input { + if request.Input[i].Content == nil { + continue + } + contentCopy := *request.Input[i].Content + if request.Input[i].Content.ContentBlocks != nil { + contentCopy.ContentBlocks = make([]schemas.ChatContentBlock, len(request.Input[i].Content.ContentBlocks)) + copy(contentCopy.ContentBlocks, request.Input[i].Content.ContentBlocks) + } + prepared.Input[i].Content = &contentCopy + } + + return &prepared +} + +func (provider *GigaChatProvider) prepareGigaChatChatAttachmentBlock(ctx *schemas.BifrostContext, key schemas.Key, blockIndex int, block schemas.ChatContentBlock) (schemas.ChatContentBlock, bool, *schemas.BifrostError) { + switch block.Type { + case schemas.ChatContentBlockTypeImage: + upload, err := gigaChatChatImageUpload(blockIndex, block) + if err != nil { + return schemas.ChatContentBlock{}, false, providerUtils.NewBifrostOperationError(err.Error(), err) + } + return provider.uploadGigaChatChatAttachment(ctx, key, upload) + case schemas.ChatContentBlockTypeFile: + if block.File == nil { + return block, false, nil + } + if block.File.FileID != nil && strings.TrimSpace(*block.File.FileID) != "" { + return block, false, nil + } + if block.File.FileURL != nil && strings.TrimSpace(*block.File.FileURL) != "" { + return schemas.ChatContentBlock{}, false, providerUtils.NewBifrostOperationError( + fmt.Sprintf("content block %d: GigaChat chat file_url attachments are not supported; upload the file first and pass file_id", blockIndex), + nil, + ) + } + if block.File.FileData == nil || strings.TrimSpace(*block.File.FileData) == "" { + return block, false, nil + } + + upload, err := gigaChatChatFileUpload(blockIndex, block.File) + if err != nil { + return schemas.ChatContentBlock{}, false, providerUtils.NewBifrostOperationError(err.Error(), err) + } + return provider.uploadGigaChatChatAttachment(ctx, key, upload) + default: + return block, false, nil + } +} + +type gigaChatChatAttachmentUpload struct { + file []byte + filename string + contentType string +} + +func (provider *GigaChatProvider) uploadGigaChatChatAttachment(ctx *schemas.BifrostContext, key schemas.Key, upload gigaChatChatAttachmentUpload) (schemas.ChatContentBlock, bool, *schemas.BifrostError) { + uploadResp, bifrostErr := provider.FileUpload(ctx, key, &schemas.BifrostFileUploadRequest{ + Provider: provider.GetProviderKey(), + File: upload.file, + Filename: upload.filename, + Purpose: schemas.FilePurposeUserData, + ContentType: &upload.contentType, + }) + if bifrostErr != nil { + return schemas.ChatContentBlock{}, false, bifrostErr + } + if uploadResp == nil || strings.TrimSpace(uploadResp.ID) == "" { + return schemas.ChatContentBlock{}, false, providerUtils.NewBifrostOperationError("GigaChat file upload response did not include file id", nil) + } + + fileID := strings.TrimSpace(uploadResp.ID) + return schemas.ChatContentBlock{ + Type: schemas.ChatContentBlockTypeFile, + File: &schemas.ChatInputFile{ + FileID: &fileID, + Filename: &upload.filename, + FileType: &upload.contentType, + }, + }, true, nil +} + +func gigaChatChatImageUpload(blockIndex int, block schemas.ChatContentBlock) (gigaChatChatAttachmentUpload, error) { + if block.ImageURLStruct == nil || strings.TrimSpace(block.ImageURLStruct.URL) == "" { + return gigaChatChatAttachmentUpload{}, fmt.Errorf("content block %d: image_url.url is required", blockIndex) + } + + sanitizedURL, err := schemas.SanitizeImageURL(block.ImageURLStruct.URL) + if err != nil { + return gigaChatChatAttachmentUpload{}, fmt.Errorf("content block %d: invalid image_url: %w", blockIndex, err) + } + urlInfo := schemas.ExtractURLTypeInfo(sanitizedURL) + if urlInfo.Type != schemas.ImageContentTypeBase64 || urlInfo.DataURLWithoutPrefix == nil { + return gigaChatChatAttachmentUpload{}, fmt.Errorf("content block %d: GigaChat chat image_url supports base64 data URLs only; upload remote images first and pass file_id", blockIndex) + } + + fileBytes, err := decodeGigaChatAttachmentBase64(*urlInfo.DataURLWithoutPrefix) + if err != nil { + return gigaChatChatAttachmentUpload{}, fmt.Errorf("content block %d: failed to decode image_url data: %w", blockIndex, err) + } + contentType := "image/jpeg" + if urlInfo.MediaType != nil && strings.TrimSpace(*urlInfo.MediaType) != "" { + contentType = normalizeGigaChatContentType(*urlInfo.MediaType) + } + + return gigaChatChatAttachmentUpload{ + file: fileBytes, + filename: "image" + extensionForGigaChatContentType(contentType), + contentType: contentType, + }, nil +} + +func gigaChatChatFileUpload(blockIndex int, file *schemas.ChatInputFile) (gigaChatChatAttachmentUpload, error) { + contentType := strings.TrimSpace(valueOrEmpty(file.FileType)) + filename := strings.TrimSpace(valueOrEmpty(file.Filename)) + fileData := strings.TrimSpace(valueOrEmpty(file.FileData)) + if contentType == "" && filename != "" { + contentType = mime.TypeByExtension(strings.ToLower(filepath.Ext(filename))) + } + + if dataURLContentType, dataURLPayload, isBase64, ok, err := parseGigaChatDataURL(fileData); err != nil { + return gigaChatChatAttachmentUpload{}, fmt.Errorf("content block %d: invalid file_data data URL: %w", blockIndex, err) + } else if ok { + if contentType == "" { + contentType = dataURLContentType + } + if isBase64 { + decoded, err := decodeGigaChatAttachmentBase64(dataURLPayload) + if err != nil { + return gigaChatChatAttachmentUpload{}, fmt.Errorf("content block %d: failed to decode file_data: %w", blockIndex, err) + } + return gigaChatChatAttachmentUpload{ + file: decoded, + filename: filenameForGigaChatAttachment(filename, contentType, "file"), + contentType: normalizeGigaChatContentType(contentType), + }, nil + } + return gigaChatChatAttachmentUpload{ + file: []byte(dataURLPayload), + filename: filenameForGigaChatAttachment(filename, contentType, "file"), + contentType: normalizeGigaChatContentType(contentType), + }, nil + } + + if isGigaChatTextContentType(contentType) { + return gigaChatChatAttachmentUpload{ + file: []byte(fileData), + filename: filenameForGigaChatAttachment(filename, contentType, "file"), + contentType: normalizeGigaChatContentType(contentType), + }, nil + } + + decoded, err := decodeGigaChatAttachmentBase64(fileData) + if err != nil { + return gigaChatChatAttachmentUpload{}, fmt.Errorf("content block %d: file_data must be a base64 data URL or base64-encoded content: %w", blockIndex, err) + } + + if contentType == "" { + contentType = "application/octet-stream" + } + return gigaChatChatAttachmentUpload{ + file: decoded, + filename: filenameForGigaChatAttachment(filename, contentType, "file"), + contentType: normalizeGigaChatContentType(contentType), + }, nil +} + +func parseGigaChatDataURL(value string) (string, string, bool, bool, error) { + if !strings.HasPrefix(value, "data:") { + return "", "", false, false, nil + } + + metaAndPayload := strings.TrimPrefix(value, "data:") + meta, payload, ok := strings.Cut(metaAndPayload, ",") + if !ok { + return "", "", false, true, fmt.Errorf("missing comma separator") + } + + contentType := "text/plain" + isBase64 := false + for index, part := range strings.Split(meta, ";") { + if index == 0 && strings.TrimSpace(part) != "" { + contentType = strings.TrimSpace(part) + continue + } + if strings.EqualFold(strings.TrimSpace(part), "base64") { + isBase64 = true + } + } + if isBase64 { + return normalizeGigaChatContentType(contentType), payload, true, true, nil + } + + decodedPayload, err := url.PathUnescape(payload) + if err != nil { + return "", "", false, true, err + } + return normalizeGigaChatContentType(contentType), decodedPayload, false, true, nil +} + +func decodeGigaChatAttachmentBase64(value string) ([]byte, error) { + cleaned := strings.Map(func(r rune) rune { + switch r { + case ' ', '\n', '\r', '\t': + return -1 + default: + return r + } + }, value) + + decoders := []*base64.Encoding{ + base64.StdEncoding, + base64.RawStdEncoding, + base64.URLEncoding, + base64.RawURLEncoding, + } + var lastErr error + for _, decoder := range decoders { + decoded, err := decoder.DecodeString(cleaned) + if err == nil { + return decoded, nil + } + lastErr = err + } + return nil, lastErr +} + +func filenameForGigaChatAttachment(filename string, contentType string, fallbackBase string) string { + if strings.TrimSpace(filename) != "" { + return strings.TrimSpace(filename) + } + return fallbackBase + extensionForGigaChatContentType(contentType) +} + +func extensionForGigaChatContentType(contentType string) string { + switch normalizeGigaChatContentType(contentType) { + case "image/jpeg": + return ".jpg" + case "application/pdf": + return ".pdf" + case "text/plain": + return ".txt" + case "application/json": + return ".json" + case "audio/mpeg": + return ".mp3" + case "audio/mp4": + return ".mp4" + } + if ext, err := mime.ExtensionsByType(normalizeGigaChatContentType(contentType)); err == nil && len(ext) > 0 { + return ext[0] + } + return "" +} + +func normalizeGigaChatContentType(contentType string) string { + contentType = strings.TrimSpace(strings.ToLower(contentType)) + switch contentType { + case "": + return "application/octet-stream" + case "image/jpg": + return "image/jpeg" + default: + return contentType + } +} + +func valueOrEmpty(value *string) string { + if value == nil { + return "" + } + return *value +} + +func isGigaChatTextContentType(contentType string) bool { + contentType = normalizeGigaChatContentType(contentType) + return strings.HasPrefix(contentType, "text/") || + contentType == "application/json" || + contentType == "application/xml" || + contentType == "application/yaml" || + contentType == "application/x-yaml" +} + +func gigaChatChatFileRequiresAutoFunctionCall(file *schemas.ChatInputFile) bool { + if file == nil { + return false + } + contentType := strings.TrimSpace(valueOrEmpty(file.FileType)) + if contentType == "" { + filename := strings.TrimSpace(valueOrEmpty(file.Filename)) + if filename != "" { + contentType = mime.TypeByExtension(strings.ToLower(filepath.Ext(filename))) + } + } + contentType = normalizeGigaChatContentType(contentType) + return contentType != "" && + contentType != "application/octet-stream" && + !strings.HasPrefix(contentType, "image/") +} diff --git a/core/providers/gigachat/count_tokens.go b/core/providers/gigachat/count_tokens.go new file mode 100644 index 00000000000..f0ad5b5595b --- /dev/null +++ b/core/providers/gigachat/count_tokens.go @@ -0,0 +1,208 @@ +package gigachat + +import ( + "fmt" + "strings" + + schemas "github.com/maximhq/bifrost/core/schemas" +) + +// ToGigaChatCountTokensRequest converts a Bifrost Responses request to the +// GigaChat /tokens/count payload. +func ToGigaChatCountTokensRequest(bifrostReq *schemas.BifrostResponsesRequest) (*GigaChatCountTokensRequest, error) { + if bifrostReq == nil { + return nil, fmt.Errorf("bifrost count tokens request is nil") + } + if strings.TrimSpace(bifrostReq.Model) == "" { + return nil, fmt.Errorf("model is required") + } + + input, err := toGigaChatCountTokensInput(bifrostReq.Input) + if err != nil { + return nil, err + } + + return &GigaChatCountTokensRequest{ + Model: bifrostReq.Model, + Input: input, + }, nil +} + +// ToBifrostCountTokensResponse converts a GigaChat /tokens/count response to the +// shared Bifrost count tokens response. +func ToBifrostCountTokensResponse(providerName schemas.ModelProvider, response *GigaChatCountTokensResponse, model string) *schemas.BifrostCountTokensResponse { + if response == nil { + return nil + } + + counts := response.Items + if len(counts) == 0 && len(response.Data) > 0 { + counts = response.Data + } + if len(counts) == 0 && response.Tokens != nil { + counts = []GigaChatCountTokensItem{{Tokens: *response.Tokens}} + } + if len(counts) == 0 { + return nil + } + + tokens := make([]int, 0, len(counts)) + inputTokens := 0 + for _, count := range counts { + tokens = append(tokens, count.Tokens) + inputTokens += count.Tokens + } + totalTokens := inputTokens + object := response.Object + if strings.TrimSpace(object) == "" { + object = "response.input_tokens" + } + if strings.TrimSpace(model) == "" { + model = response.Model + } + + return &schemas.BifrostCountTokensResponse{ + Object: object, + Model: model, + InputTokens: inputTokens, + InputTokensDetails: &schemas.ResponsesResponseInputTokens{ + TextTokens: inputTokens, + }, + Tokens: tokens, + TotalTokens: &totalTokens, + ExtraFields: schemas.BifrostResponseExtraFields{ + Provider: providerName, + }, + } +} + +func toGigaChatCountTokensInput(messages []schemas.ResponsesMessage) ([]string, error) { + if len(messages) == 0 { + return nil, fmt.Errorf("count tokens input is required") + } + + input := make([]string, 0, len(messages)) + for index, message := range messages { + parts, err := toGigaChatCountTokensMessageInput(index, message) + if err != nil { + return nil, err + } + input = append(input, parts...) + } + if len(input) == 0 { + return nil, fmt.Errorf("count tokens text input is empty after conversion") + } + return input, nil +} + +func toGigaChatCountTokensMessageInput(index int, message schemas.ResponsesMessage) ([]string, error) { + parts := make([]string, 0, 1) + if message.Content != nil { + contentParts, err := toGigaChatCountTokensContentInput(index, message.Content) + if err != nil { + return nil, err + } + parts = append(parts, contentParts...) + } + + if message.ResponsesReasoning != nil { + if message.ResponsesReasoning.EncryptedContent != nil && strings.TrimSpace(*message.ResponsesReasoning.EncryptedContent) != "" { + return nil, fmt.Errorf("GigaChat count tokens supports only text input; input[%d] contains encrypted reasoning content", index) + } + for _, summary := range message.ResponsesReasoning.Summary { + if text := strings.TrimSpace(summary.Text); text != "" { + parts = append(parts, summary.Text) + } + } + } + + if len(parts) == 0 && hasGigaChatCountTokensNonTextMessagePayload(message) { + return nil, fmt.Errorf("GigaChat count tokens supports only text input; input[%d] contains a non-text item", index) + } + return parts, nil +} + +func toGigaChatCountTokensContentInput(index int, content *schemas.ResponsesMessageContent) ([]string, error) { + if content.ContentStr != nil && content.ContentBlocks != nil { + return nil, fmt.Errorf("input[%d].content cannot contain both string content and content blocks", index) + } + + if content.ContentStr != nil { + if strings.TrimSpace(*content.ContentStr) == "" { + return nil, nil + } + return []string{*content.ContentStr}, nil + } + + parts := make([]string, 0, len(content.ContentBlocks)) + for blockIndex, block := range content.ContentBlocks { + text, ok, err := toGigaChatCountTokensContentBlockText(index, blockIndex, block) + if err != nil { + return nil, err + } + if ok && strings.TrimSpace(text) != "" { + parts = append(parts, text) + } + } + return parts, nil +} + +func toGigaChatCountTokensContentBlockText(index int, blockIndex int, block schemas.ResponsesMessageContentBlock) (string, bool, error) { + if isGigaChatCountTokensUnsupportedMediaBlock(block) { + return "", false, fmt.Errorf("GigaChat count tokens supports only text input; input[%d].content[%d] contains file, image, or audio content", index, blockIndex) + } + + if block.ResponsesOutputMessageContentText != nil { + if block.Text != nil && strings.TrimSpace(*block.Text) != "" { + return *block.Text, true, nil + } + return "", false, nil + } + if block.Text != nil { + if strings.TrimSpace(*block.Text) == "" { + return "", false, nil + } + return *block.Text, true, nil + } + if block.ResponsesOutputMessageContentRefusal != nil && strings.TrimSpace(block.ResponsesOutputMessageContentRefusal.Refusal) != "" { + return block.ResponsesOutputMessageContentRefusal.Refusal, true, nil + } + + switch block.Type { + case schemas.ResponsesInputMessageContentBlockTypeText, + schemas.ResponsesOutputMessageContentTypeText, + schemas.ResponsesOutputMessageContentTypeReasoning, + schemas.ResponsesOutputMessageContentTypeRefusal: + return "", false, nil + case "": + if !hasGigaChatCountTokensNonTextBlockPayload(block) { + return "", false, nil + } + } + + return "", false, fmt.Errorf("GigaChat count tokens supports only text input; input[%d].content[%d] has unsupported content type %q", index, blockIndex, block.Type) +} + +func isGigaChatCountTokensUnsupportedMediaBlock(block schemas.ResponsesMessageContentBlock) bool { + return block.FileID != nil || + block.ResponsesInputMessageContentBlockImage != nil || + block.ResponsesInputMessageContentBlockFile != nil || + block.Audio != nil || + block.Type == schemas.ResponsesInputMessageContentBlockTypeImage || + block.Type == schemas.ResponsesInputMessageContentBlockTypeFile || + block.Type == schemas.ResponsesInputMessageContentBlockTypeAudio +} + +func hasGigaChatCountTokensNonTextBlockPayload(block schemas.ResponsesMessageContentBlock) bool { + return block.Signature != nil || + block.ResponsesOutputMessageContentRenderedContent != nil || + block.ResponsesOutputMessageContentCompaction != nil || + block.CacheControl != nil || + block.Citations != nil +} + +func hasGigaChatCountTokensNonTextMessagePayload(message schemas.ResponsesMessage) bool { + return message.ResponsesToolMessage != nil || + message.CacheControl != nil || + message.Type != nil && *message.Type != schemas.ResponsesMessageTypeMessage +} diff --git a/core/providers/gigachat/embedding.go b/core/providers/gigachat/embedding.go new file mode 100644 index 00000000000..88919fbe5d0 --- /dev/null +++ b/core/providers/gigachat/embedding.go @@ -0,0 +1,170 @@ +package gigachat + +import ( + "encoding/base64" + "encoding/binary" + "fmt" + "math" + "sort" + "strings" + + schemas "github.com/maximhq/bifrost/core/schemas" +) + +// ToGigaChatEmbeddingRequest converts a Bifrost embedding request to GigaChat v1 format. +func ToGigaChatEmbeddingRequest(bifrostReq *schemas.BifrostEmbeddingRequest) (*GigaChatEmbeddingRequest, error) { + if bifrostReq == nil { + return nil, fmt.Errorf("bifrost embedding request is nil") + } + if strings.TrimSpace(bifrostReq.Model) == "" { + return nil, fmt.Errorf("model is required") + } + if bifrostReq.Input == nil { + return nil, fmt.Errorf("input is required") + } + if bifrostReq.Input.Text == nil && bifrostReq.Input.Texts == nil { + return nil, fmt.Errorf("GigaChat embeddings support only string or array-of-string input") + } + + if err := validateGigaChatEmbeddingEncodingFormat(bifrostReq.Params); err != nil { + return nil, err + } + if unsupportedParams := unsupportedGigaChatEmbeddingParams(bifrostReq.Params); len(unsupportedParams) > 0 { + return nil, fmt.Errorf("GigaChat embeddings do not support parameter(s): %s", strings.Join(unsupportedParams, ", ")) + } + + return &GigaChatEmbeddingRequest{ + Model: bifrostReq.Model, + Input: bifrostReq.Input, + }, nil +} + +// ToBifrostEmbeddingResponse converts a GigaChat v1 embeddings response to Bifrost format. +func ToBifrostEmbeddingResponse(providerName schemas.ModelProvider, response *GigaChatEmbeddingResponse) *schemas.BifrostEmbeddingResponse { + if response == nil { + return nil + } + + data := make([]schemas.EmbeddingData, 0, len(response.Data)) + itemUsage := &GigaChatEmbeddingUsage{} + hasItemUsage := false + for _, item := range response.Data { + embedding := append([]float64(nil), item.Embedding...) + object := item.Object + if object == "" { + object = "embedding" + } + data = append(data, schemas.EmbeddingData{ + Index: item.Index, + Object: object, + Embedding: schemas.EmbeddingStruct{ + EmbeddingArray: embedding, + }, + }) + if item.Usage != nil { + hasItemUsage = true + itemUsage.PromptTokens += item.Usage.PromptTokens + itemUsage.TotalTokens += item.Usage.TotalTokens + } + } + + object := response.Object + if object == "" { + object = "list" + } + + bifrostResponse := &schemas.BifrostEmbeddingResponse{ + Data: data, + Model: response.Model, + Object: object, + Usage: toBifrostGigaChatEmbeddingUsage(response.Usage), + ExtraFields: schemas.BifrostResponseExtraFields{ + Provider: providerName, + }, + } + if bifrostResponse.Usage == nil && hasItemUsage { + bifrostResponse.Usage = toBifrostGigaChatEmbeddingUsage(itemUsage) + } + return bifrostResponse +} + +func validateGigaChatEmbeddingEncodingFormat(params *schemas.EmbeddingParameters) error { + format := normalizedGigaChatEmbeddingEncodingFormat(params) + switch format { + case "", "float", "base64": + return nil + default: + return fmt.Errorf("GigaChat embeddings do not support encoding_format %q", *params.EncodingFormat) + } +} + +func applyGigaChatEmbeddingEncodingFormat(response *schemas.BifrostEmbeddingResponse, params *schemas.EmbeddingParameters) error { + if normalizedGigaChatEmbeddingEncodingFormat(params) != "base64" { + return nil + } + if response == nil { + return nil + } + + for i := range response.Data { + if response.Data[i].Embedding.EmbeddingStr != nil { + continue + } + if response.Data[i].Embedding.EmbeddingArray == nil { + return fmt.Errorf("GigaChat embeddings cannot encode non-float embedding at index %d as base64", i) + } + + encoded := encodeGigaChatEmbeddingFloat32Base64(response.Data[i].Embedding.EmbeddingArray) + response.Data[i].Embedding = schemas.EmbeddingStruct{EmbeddingStr: &encoded} + } + return nil +} + +func normalizedGigaChatEmbeddingEncodingFormat(params *schemas.EmbeddingParameters) string { + if params == nil || params.EncodingFormat == nil { + return "" + } + return strings.ToLower(strings.TrimSpace(*params.EncodingFormat)) +} + +func encodeGigaChatEmbeddingFloat32Base64(values []float64) string { + buf := make([]byte, len(values)*4) + for i, value := range values { + binary.LittleEndian.PutUint32(buf[i*4:], math.Float32bits(float32(value))) + } + return base64.StdEncoding.EncodeToString(buf) +} + +func unsupportedGigaChatEmbeddingParams(params *schemas.EmbeddingParameters) []string { + if params == nil { + return nil + } + + unsupported := make([]string, 0) + if params.Dimensions != nil { + unsupported = append(unsupported, "dimensions") + } + for name := range params.ExtraParams { + if strings.TrimSpace(name) != "" { + unsupported = append(unsupported, name) + } + } + + sort.Strings(unsupported) + return unsupported +} + +func toBifrostGigaChatEmbeddingUsage(usage *GigaChatEmbeddingUsage) *schemas.BifrostLLMUsage { + if usage == nil { + return nil + } + totalTokens := usage.TotalTokens + if totalTokens == 0 && usage.PromptTokens > 0 { + totalTokens = usage.PromptTokens + } + return &schemas.BifrostLLMUsage{ + PromptTokens: usage.PromptTokens, + CompletionTokens: 0, + TotalTokens: totalTokens, + } +} diff --git a/core/providers/gigachat/errors.go b/core/providers/gigachat/errors.go new file mode 100644 index 00000000000..cc67e9ad552 --- /dev/null +++ b/core/providers/gigachat/errors.go @@ -0,0 +1,76 @@ +package gigachat + +import ( + "bytes" + "encoding/json" + "fmt" + "strconv" + "strings" + + providerUtils "github.com/maximhq/bifrost/core/providers/utils" + schemas "github.com/maximhq/bifrost/core/schemas" + "github.com/valyala/fasthttp" +) + +// ParseGigaChatError parses GigaChat REST error responses. +func ParseGigaChatError(resp *fasthttp.Response, providerName schemas.ModelProvider) *schemas.BifrostError { + var errorResp GigaChatErrorResponse + bifrostErr := providerUtils.HandleProviderAPIError(resp, &errorResp) + if bifrostErr.Error == nil { + bifrostErr.Error = &schemas.ErrorField{} + } + + if message := gigaChatErrorMessage(errorResp); message != "" { + bifrostErr.Error.Message = message + } + if codeValue, ok := gigaChatErrorCode(errorResp); ok { + code := codeValue + bifrostErr.Error.Code = &code + } else if errorResp.Status != nil { + code := strconv.Itoa(*errorResp.Status) + bifrostErr.Error.Code = &code + } + if errorResp.Status != nil && bifrostErr.StatusCode == nil { + status := *errorResp.Status + bifrostErr.StatusCode = &status + } + if strings.TrimSpace(bifrostErr.Error.Message) == "" { + if bifrostErr.StatusCode != nil { + bifrostErr.Error.Message = fmt.Sprintf("GigaChat API error (status %d)", *bifrostErr.StatusCode) + } else { + bifrostErr.Error.Message = "GigaChat API error" + } + } + bifrostErr.Error.Message = redactGigaChatSensitiveText(bifrostErr.Error.Message) + + bifrostErr.ExtraFields.Provider = providerName + return bifrostErr +} + +func gigaChatErrorMessage(errorResp GigaChatErrorResponse) string { + for _, message := range []string{errorResp.Message, errorResp.ErrorDescription, errorResp.Error} { + if trimmed := strings.TrimSpace(message); trimmed != "" { + return trimmed + } + } + return "" +} + +func gigaChatErrorCode(errorResp GigaChatErrorResponse) (string, bool) { + code := bytes.TrimSpace(errorResp.Code) + if len(code) == 0 || bytes.Equal(code, []byte("null")) { + return "", false + } + + if len(code) >= 2 && code[0] == '"' { + var value string + if err := json.Unmarshal(code, &value); err != nil { + return "", false + } + trimmed := strings.TrimSpace(value) + return trimmed, trimmed != "" + } + + trimmed := strings.TrimSpace(string(code)) + return trimmed, trimmed != "" +} diff --git a/core/providers/gigachat/files.go b/core/providers/gigachat/files.go new file mode 100644 index 00000000000..f89389e557b --- /dev/null +++ b/core/providers/gigachat/files.go @@ -0,0 +1,617 @@ +package gigachat + +import ( + "bytes" + "encoding/base64" + "encoding/json" + "fmt" + "mime" + "mime/multipart" + "net/http" + "net/textproto" + "net/url" + "path/filepath" + "strconv" + "strings" + "time" + "unicode" + "unicode/utf8" + + providerUtils "github.com/maximhq/bifrost/core/providers/utils" + schemas "github.com/maximhq/bifrost/core/schemas" + "github.com/valyala/fasthttp" +) + +const ( + gigaChatFilePurposeAssistant = "assistant" + gigaChatFilePurposeGeneral = "general" +) + +func emptyGigaChatJSONRequestBody() []byte { + return []byte("{}") +} + +func toGigaChatFilePurpose(purpose schemas.FilePurpose) string { + if purpose == schemas.FilePurposeAssistants { + return gigaChatFilePurposeAssistant + } + return gigaChatFilePurposeGeneral +} + +func toBifrostFilePurpose(gigaChatPurpose string, requestedPurpose schemas.FilePurpose) schemas.FilePurpose { + switch strings.TrimSpace(gigaChatPurpose) { + case gigaChatFilePurposeAssistant: + return schemas.FilePurposeAssistants + case gigaChatFilePurposeGeneral, "": + if requestedPurpose != "" { + return requestedPurpose + } + return schemas.FilePurposeUserData + default: + return schemas.FilePurpose(gigaChatPurpose) + } +} + +func paginateGigaChatFileList(response *schemas.BifrostFileListResponse, nativeCursor string, limit int) (string, bool, error) { + offset := 0 + if nativeCursor != "" { + parsed, err := strconv.Atoi(nativeCursor) + if err != nil || parsed < 0 { + return "", false, fmt.Errorf("invalid GigaChat file cursor %q", nativeCursor) + } + offset = parsed + } + if offset > len(response.Data) { + return "", false, fmt.Errorf("GigaChat file cursor offset %d exceeds result count %d", offset, len(response.Data)) + } + + end := len(response.Data) + hasMore := false + if limit > 0 && limit < len(response.Data)-offset { + end = offset + limit + hasMore = true + } + response.Data = response.Data[offset:end] + if !hasMore { + return "", false, nil + } + return strconv.Itoa(end), true, nil +} + +func (provider *GigaChatProvider) fileUploadWithRefresh(ctx *schemas.BifrostContext, key schemas.Key, request *schemas.BifrostFileUploadRequest, forceRefresh bool) (*schemas.BifrostFileUploadResponse, *schemas.BifrostError) { + ctx = ensureGigaChatContext(ctx) + + body, contentType, bifrostErr := buildGigaChatFileUploadBody(request) + if bifrostErr != nil { + return nil, bifrostErr + } + + responseBody, providerResponseHeaders, _, latency, bifrostErr := provider.executeGigaChatFileRequest(ctx, key, schemas.FileUploadRequest, http.MethodPost, "/files", contentType, "application/json", body, nil, forceRefresh) + if bifrostErr != nil { + return nil, bifrostErr + } + + sendBackRawRequest := providerUtils.ShouldSendBackRawRequest(ctx, provider.sendBackRawRequest) + sendBackRawResponse := providerUtils.ShouldSendBackRawResponse(ctx, provider.sendBackRawResponse) + + gigaChatResponse := &GigaChatUploadedFile{} + rawRequest, rawResponse, bifrostErr := providerUtils.HandleProviderResponse(responseBody, gigaChatResponse, nil, sendBackRawRequest, sendBackRawResponse) + if bifrostErr != nil { + return nil, enrichGigaChatError(ctx, bifrostErr, nil, responseBody, sendBackRawRequest, sendBackRawResponse) + } + + response := toBifrostFileUploadResponse(*gigaChatResponse, request.Purpose, latency) + response.ExtraFields.ProviderResponseHeaders = providerResponseHeaders + if sendBackRawRequest { + response.ExtraFields.RawRequest = rawRequest + } + if sendBackRawResponse { + response.ExtraFields.RawResponse = rawResponse + } + return response, nil +} + +func (provider *GigaChatProvider) fileListWithRefresh(ctx *schemas.BifrostContext, key schemas.Key, request *schemas.BifrostFileListRequest, forceRefresh bool) (*schemas.BifrostFileListResponse, *schemas.BifrostError) { + ctx = ensureGigaChatContext(ctx) + + rawRequestBody := emptyGigaChatJSONRequestBody() + responseBody, providerResponseHeaders, _, latency, bifrostErr := provider.executeGigaChatFileRequest(ctx, key, schemas.FileListRequest, http.MethodGet, "/files", "", "application/json", nil, rawRequestBody, forceRefresh) + if bifrostErr != nil { + return nil, bifrostErr + } + + sendBackRawRequest := providerUtils.ShouldSendBackRawRequest(ctx, provider.sendBackRawRequest) + sendBackRawResponse := providerUtils.ShouldSendBackRawResponse(ctx, provider.sendBackRawResponse) + + gigaChatResponse := &GigaChatUploadedFiles{} + rawRequest, rawResponse, bifrostErr := providerUtils.HandleProviderResponse(responseBody, gigaChatResponse, rawRequestBody, sendBackRawRequest, sendBackRawResponse) + if bifrostErr != nil { + return nil, enrichGigaChatError(ctx, bifrostErr, rawRequestBody, responseBody, sendBackRawRequest, sendBackRawResponse) + } + + files := make([]schemas.FileObject, 0, len(gigaChatResponse.Data)) + for _, file := range gigaChatResponse.Data { + converted := toBifrostFileObject(file, "") + if request != nil && request.Purpose != "" && converted.Purpose != request.Purpose { + continue + } + files = append(files, converted) + } + + response := &schemas.BifrostFileListResponse{ + Object: "list", + Data: files, + ExtraFields: schemas.BifrostResponseExtraFields{ + Latency: latency.Milliseconds(), + ProviderResponseHeaders: providerResponseHeaders, + }, + } + if sendBackRawRequest { + response.ExtraFields.RawRequest = rawRequest + } + if sendBackRawResponse { + response.ExtraFields.RawResponse = rawResponse + } + return response, nil +} + +func (provider *GigaChatProvider) fileRetrieveWithRefresh(ctx *schemas.BifrostContext, key schemas.Key, request *schemas.BifrostFileRetrieveRequest, forceRefresh bool) (*schemas.BifrostFileRetrieveResponse, *schemas.BifrostError) { + ctx = ensureGigaChatContext(ctx) + + path := fmt.Sprintf("/files/%s", url.PathEscape(request.FileID)) + responseBody, providerResponseHeaders, _, latency, bifrostErr := provider.executeGigaChatFileRequest(ctx, key, schemas.FileRetrieveRequest, http.MethodGet, path, "", "application/json", nil, nil, forceRefresh) + if bifrostErr != nil { + return nil, bifrostErr + } + + sendBackRawRequest := providerUtils.ShouldSendBackRawRequest(ctx, provider.sendBackRawRequest) + sendBackRawResponse := providerUtils.ShouldSendBackRawResponse(ctx, provider.sendBackRawResponse) + + gigaChatResponse := &GigaChatUploadedFile{} + rawRequest, rawResponse, bifrostErr := providerUtils.HandleProviderResponse(responseBody, gigaChatResponse, nil, sendBackRawRequest, sendBackRawResponse) + if bifrostErr != nil { + return nil, enrichGigaChatError(ctx, bifrostErr, nil, responseBody, sendBackRawRequest, sendBackRawResponse) + } + + response := toBifrostFileRetrieveResponse(*gigaChatResponse, "", latency) + response.ExtraFields.ProviderResponseHeaders = providerResponseHeaders + if sendBackRawRequest { + response.ExtraFields.RawRequest = rawRequest + } + if sendBackRawResponse { + response.ExtraFields.RawResponse = rawResponse + } + return response, nil +} + +func (provider *GigaChatProvider) fileDeleteWithRefresh(ctx *schemas.BifrostContext, key schemas.Key, request *schemas.BifrostFileDeleteRequest, forceRefresh bool) (*schemas.BifrostFileDeleteResponse, *schemas.BifrostError) { + ctx = ensureGigaChatContext(ctx) + + path := fmt.Sprintf("/files/%s/delete", url.PathEscape(request.FileID)) + responseBody, providerResponseHeaders, _, latency, bifrostErr := provider.executeGigaChatFileRequest(ctx, key, schemas.FileDeleteRequest, http.MethodPost, path, "", "application/json", nil, nil, forceRefresh) + if bifrostErr != nil { + return nil, bifrostErr + } + + sendBackRawRequest := providerUtils.ShouldSendBackRawRequest(ctx, provider.sendBackRawRequest) + sendBackRawResponse := providerUtils.ShouldSendBackRawResponse(ctx, provider.sendBackRawResponse) + + gigaChatResponse := &GigaChatDeletedFile{} + rawRequest, rawResponse, bifrostErr := providerUtils.HandleProviderResponse(responseBody, gigaChatResponse, nil, sendBackRawRequest, sendBackRawResponse) + if bifrostErr != nil { + return nil, enrichGigaChatError(ctx, bifrostErr, nil, responseBody, sendBackRawRequest, sendBackRawResponse) + } + + response := &schemas.BifrostFileDeleteResponse{ + ID: gigaChatResponse.ID, + Object: "file", + Deleted: gigaChatResponse.Deleted, + ExtraFields: schemas.BifrostResponseExtraFields{ + Latency: latency.Milliseconds(), + ProviderResponseHeaders: providerResponseHeaders, + }, + } + if sendBackRawRequest { + response.ExtraFields.RawRequest = rawRequest + } + if sendBackRawResponse { + response.ExtraFields.RawResponse = rawResponse + } + return response, nil +} + +func (provider *GigaChatProvider) fileContentWithRefresh(ctx *schemas.BifrostContext, key schemas.Key, request *schemas.BifrostFileContentRequest, forceRefresh bool) (*schemas.BifrostFileContentResponse, *schemas.BifrostError) { + ctx = ensureGigaChatContext(ctx) + + path := fmt.Sprintf("/files/%s/content", url.PathEscape(request.FileID)) + responseBody, providerResponseHeaders, responseContentType, latency, bifrostErr := provider.executeGigaChatFileRequest(ctx, key, schemas.FileContentRequest, http.MethodGet, path, "", "*/*", nil, nil, forceRefresh) + if bifrostErr != nil { + return nil, bifrostErr + } + + contentType := responseContentType + content, decodedContentType, decodeErr := decodeGigaChatFileContent(responseBody, contentType) + if decodeErr != nil { + return nil, newGigaChatProviderResponseError("failed to decode GigaChat file content response", decodeErr) + } + if decodedContentType != "" { + contentType = decodedContentType + } + if strings.TrimSpace(contentType) == "" { + contentType = "application/octet-stream" + } + + return &schemas.BifrostFileContentResponse{ + FileID: request.FileID, + Content: content, + ContentType: contentType, + ExtraFields: schemas.BifrostResponseExtraFields{ + Latency: latency.Milliseconds(), + ProviderResponseHeaders: providerResponseHeaders, + }, + }, nil +} + +func (provider *GigaChatProvider) executeGigaChatFileRequest( + ctx *schemas.BifrostContext, + key schemas.Key, + requestType schemas.RequestType, + method string, + path string, + contentType string, + accept string, + body []byte, + rawRequestForError []byte, + forceRefresh bool, +) ([]byte, map[string]string, string, time.Duration, *schemas.BifrostError) { + headers, bifrostErr := provider.buildAuthHeadersWithRefresh(ctx, key, forceRefresh) + if bifrostErr != nil { + return nil, nil, "", 0, bifrostErr + } + + client, clientErr := provider.getGigaChatTLSClient(provider.client, gigaChatTLSClientCacheDefault, key.GigaChatKeyConfig) + if clientErr != nil { + return nil, nil, "", 0, newGigaChatConfigurationError(clientErr.Error()) + } + + req := fasthttp.AcquireRequest() + resp := fasthttp.AcquireResponse() + defer fasthttp.ReleaseRequest(req) + defer fasthttp.ReleaseResponse(resp) + + for headerName, headerValue := range headers { + req.Header.Set(headerName, headerValue) + } + req.SetRequestURI(buildGigaChatRequestURL(ctx, resolveBaseURL(key, provider.networkConfig), gigaChatAPIVersionV1, path, provider.customProviderConfig, requestType)) + req.Header.SetMethod(method) + if strings.TrimSpace(contentType) != "" { + req.Header.SetContentType(contentType) + } + if strings.TrimSpace(accept) != "" { + req.Header.Set("Accept", accept) + } + if body != nil { + req.SetBody(body) + } + + sendBackRawRequest := providerUtils.ShouldSendBackRawRequest(ctx, provider.sendBackRawRequest) + sendBackRawResponse := providerUtils.ShouldSendBackRawResponse(ctx, provider.sendBackRawResponse) + + latency, bifrostErr, wait := providerUtils.MakeRequestWithContext(ctx, client, req, resp) + defer wait() + if bifrostErr != nil { + bifrostErr.ExtraFields.Provider = provider.GetProviderKey() + return nil, nil, "", latency, enrichGigaChatError(ctx, bifrostErr, rawRequestForError, nil, sendBackRawRequest, sendBackRawResponse) + } + + responseContentType := string(resp.Header.ContentType()) + providerResponseHeaders := providerUtils.ExtractProviderResponseHeaders(resp) + ctx.SetValue(schemas.BifrostContextKeyProviderResponseHeaders, providerResponseHeaders) + + if resp.StatusCode() < fasthttp.StatusOK || resp.StatusCode() >= fasthttp.StatusMultipleChoices { + bifrostErr := ParseGigaChatError(resp, provider.GetProviderKey()) + return nil, providerResponseHeaders, responseContentType, latency, enrichGigaChatError(ctx, bifrostErr, rawRequestForError, resp.Body(), sendBackRawRequest, sendBackRawResponse) + } + + responseBody, err := providerUtils.CheckAndDecodeBody(resp) + if err != nil { + bifrostErr := newGigaChatProviderResponseError("failed to decode GigaChat file response", err) + return nil, providerResponseHeaders, responseContentType, latency, enrichGigaChatError(ctx, bifrostErr, rawRequestForError, resp.Body(), sendBackRawRequest, sendBackRawResponse) + } + + return responseBody, providerResponseHeaders, responseContentType, latency, nil +} + +func buildGigaChatFileUploadBody(request *schemas.BifrostFileUploadRequest) ([]byte, string, *schemas.BifrostError) { + if request == nil { + return nil, "", providerUtils.NewBifrostOperationError("file upload request is nil", nil) + } + if len(request.File) == 0 { + return nil, "", providerUtils.NewBifrostOperationError("file content is required", nil) + } + if request.Purpose == "" { + return nil, "", providerUtils.NewBifrostOperationError("purpose is required", nil) + } + + var body bytes.Buffer + writer := multipart.NewWriter(&body) + + if err := writer.WriteField("purpose", toGigaChatFilePurpose(request.Purpose)); err != nil { + return nil, "", providerUtils.NewBifrostOperationError("failed to write purpose field", err) + } + + contentType := resolveGigaChatFileUploadContentType(request.ContentType, request.Filename, request.File) + filename := resolveGigaChatFileUploadFilename(request.Filename, contentType) + + partHeaders := textproto.MIMEHeader{} + partHeaders.Set("Content-Disposition", fmt.Sprintf(`form-data; name="file"; filename="%s"`, escapeGigaChatMultipartFilename(filename))) + if contentType != "" { + partHeaders.Set("Content-Type", contentType) + } + part, err := writer.CreatePart(partHeaders) + if err != nil { + return nil, "", providerUtils.NewBifrostOperationError("failed to create form file", err) + } + if _, err := part.Write(request.File); err != nil { + return nil, "", providerUtils.NewBifrostOperationError("failed to write file content", err) + } + if err := writer.Close(); err != nil { + return nil, "", providerUtils.NewBifrostOperationError("failed to close multipart writer", err) + } + + return body.Bytes(), writer.FormDataContentType(), nil +} + +var gigaChatFileUploadContentTypesByExtension = map[string]string{ + ".txt": "text/plain", + ".doc": "application/msword", + ".docx": "application/vnd.openxmlformats-officedocument.wordprocessingml.document", + ".pdf": "application/pdf", + ".epub": "application/epub", + ".ppt": "application/ppt", + ".pptx": "application/pptx", + ".xlsx": "application/vnd.ms-excel", + ".jpg": "image/jpeg", + ".jpeg": "image/jpeg", + ".png": "image/png", + ".tif": "image/tiff", + ".tiff": "image/tiff", + ".bmp": "image/bmp", + ".mp4": "audio/mp4", + ".mp3": "audio/mp3", + ".m4a": "audio/x-m4a", + ".wav": "audio/x-wav", + ".weba": "audio/webm", + ".ogg": "audio/x-ogg", + ".opus": "audio/opus", +} + +var gigaChatFileUploadExtensionsByContentType = map[string]string{ + "text/plain": "txt", + "application/msword": "doc", + "application/vnd.openxmlformats-officedocument.wordprocessingml.document": "docx", + "application/pdf": "pdf", + "application/epub": "epub", + "application/ppt": "ppt", + "application/pptx": "pptx", + "application/vnd.ms-excel": "xlsx", + "image/jpeg": "jpg", + "image/png": "png", + "image/tiff": "tiff", + "image/bmp": "bmp", + "audio/mp4": "mp4", + "audio/mp3": "mp3", + "audio/x-m4a": "m4a", + "audio/x-wav": "wav", + "audio/webm": "weba", + "audio/x-ogg": "ogg", + "audio/opus": "opus", +} + +func resolveGigaChatFileUploadContentType(contentType *string, filename string, file []byte) string { + if contentType != nil { + if normalized := normalizeGigaChatFileUploadContentType(*contentType); normalized != "" { + return normalized + } + } + if inferred := inferGigaChatFileUploadContentTypeFromFilename(filename); inferred != "" { + return inferred + } + if detected := normalizeGigaChatFileUploadContentType(http.DetectContentType(file)); detected != "" { + return detected + } + if looksLikeTextFile(file) { + return "text/plain" + } + return "application/octet-stream" +} + +func resolveGigaChatFileUploadFilename(filename string, contentType string) string { + filename = strings.TrimSpace(filename) + if filename == "" { + filename = "file" + } + + ext := strings.ToLower(filepath.Ext(filename)) + if ext != "" { + if expectedContentType := gigaChatFileUploadContentTypesByExtension[ext]; expectedContentType == contentType { + return filename + } + } + + if extension, ok := gigaChatFileUploadExtensionsByContentType[contentType]; ok { + base := filename + if ext != "" { + base = strings.TrimSuffix(filename, filepath.Ext(filename)) + } + if strings.TrimSpace(base) == "" { + base = "file" + } + return base + "." + extension + } + + return filename +} + +func inferGigaChatFileUploadContentTypeFromFilename(filename string) string { + ext := strings.ToLower(filepath.Ext(strings.TrimSpace(filename))) + if ext == "" { + return "" + } + if contentType, ok := gigaChatFileUploadContentTypesByExtension[ext]; ok { + return contentType + } + return normalizeGigaChatFileUploadContentType(mime.TypeByExtension(ext)) +} + +func normalizeGigaChatFileUploadContentType(contentType string) string { + contentType = strings.TrimSpace(strings.ToLower(contentType)) + if contentType == "" { + return "" + } + if mediaType, _, err := mime.ParseMediaType(contentType); err == nil { + contentType = strings.TrimSpace(strings.ToLower(mediaType)) + } + + switch contentType { + case "application/json", "application/jsonl", "application/x-jsonl", "application/x-ndjson", "application/ndjson", "application/json-lines", "text/json", "text/x-jsonl", "text/markdown": + return "text/plain" + case "image/jpg": + return "image/jpeg" + case "application/epub+zip": + return "application/epub" + case "application/vnd.ms-powerpoint": + return "application/ppt" + case "application/vnd.openxmlformats-officedocument.presentationml.presentation": + return "application/pptx" + case "application/vnd.openxmlformats-officedocument.spreadsheetml.sheet": + return "application/vnd.ms-excel" + case "audio/mpeg": + return "audio/mp3" + case "audio/ogg": + return "audio/x-ogg" + case "audio/wav", "audio/wave", "audio/x-pn-wav": + return "audio/x-wav" + } + + for _, supportedContentType := range gigaChatFileUploadContentTypesByExtension { + if contentType == supportedContentType { + return contentType + } + } + return "" +} + +func looksLikeTextFile(file []byte) bool { + if len(file) == 0 { + return false + } + sample := file + if len(sample) > 512 { + sample = sample[:512] + } + return bytes.IndexByte(sample, 0) == -1 && utf8.Valid(sample) +} + +func escapeGigaChatMultipartFilename(filename string) string { + var sanitized strings.Builder + sanitized.Grow(len(filename)) + for _, r := range filename { + if unicode.IsControl(r) { + sanitized.WriteByte('_') + continue + } + sanitized.WriteRune(r) + } + return strings.NewReplacer("\\", "\\\\", `"`, "\\\"").Replace(sanitized.String()) +} + +func toBifrostFileObject(file GigaChatUploadedFile, requestedPurpose schemas.FilePurpose) schemas.FileObject { + object := file.Object + if object == "" { + object = "file" + } + return schemas.FileObject{ + ID: file.ID, + Object: object, + Bytes: file.Bytes, + CreatedAt: file.CreatedAt, + Filename: file.Filename, + Purpose: toBifrostFilePurpose(file.Purpose, requestedPurpose), + } +} + +func toBifrostFileUploadResponse(file GigaChatUploadedFile, requestedPurpose schemas.FilePurpose, latency time.Duration) *schemas.BifrostFileUploadResponse { + object := file.Object + if object == "" { + object = "file" + } + return &schemas.BifrostFileUploadResponse{ + ID: file.ID, + Object: object, + Bytes: file.Bytes, + CreatedAt: file.CreatedAt, + Filename: file.Filename, + Purpose: toBifrostFilePurpose(file.Purpose, requestedPurpose), + StorageBackend: schemas.FileStorageAPI, + ExtraFields: schemas.BifrostResponseExtraFields{ + Latency: latency.Milliseconds(), + }, + } +} + +func toBifrostFileRetrieveResponse(file GigaChatUploadedFile, requestedPurpose schemas.FilePurpose, latency time.Duration) *schemas.BifrostFileRetrieveResponse { + object := file.Object + if object == "" { + object = "file" + } + return &schemas.BifrostFileRetrieveResponse{ + ID: file.ID, + Object: object, + Bytes: file.Bytes, + CreatedAt: file.CreatedAt, + Filename: file.Filename, + Purpose: toBifrostFilePurpose(file.Purpose, requestedPurpose), + StorageBackend: schemas.FileStorageAPI, + ExtraFields: schemas.BifrostResponseExtraFields{ + Latency: latency.Milliseconds(), + }, + } +} + +func decodeGigaChatFileContent(body []byte, contentType string) ([]byte, string, error) { + if isGigaChatJSONContentType(contentType) { + var wrapper map[string]json.RawMessage + if err := json.Unmarshal(body, &wrapper); err == nil { + if rawContent, ok := wrapper["content"]; ok { + var encoded string + if err := json.Unmarshal(rawContent, &encoded); err != nil { + return nil, "", err + } + decoded, err := decodeGigaChatBase64Content(encoded) + if err != nil { + return nil, "", err + } + return decoded, "application/octet-stream", nil + } + } + } + return append([]byte(nil), body...), contentType, nil +} + +func isGigaChatJSONContentType(contentType string) bool { + contentType = strings.ToLower(strings.TrimSpace(contentType)) + return contentType == "application/json" || strings.HasPrefix(contentType, "application/json;") +} + +func decodeGigaChatBase64Content(encoded string) ([]byte, error) { + decoded, err := base64.StdEncoding.DecodeString(encoded) + if err == nil { + return decoded, nil + } + decoded, rawErr := base64.RawStdEncoding.DecodeString(encoded) + if rawErr == nil { + return decoded, nil + } + return nil, err +} diff --git a/core/providers/gigachat/gigachat.go b/core/providers/gigachat/gigachat.go new file mode 100644 index 00000000000..2db5bf47082 --- /dev/null +++ b/core/providers/gigachat/gigachat.go @@ -0,0 +1,1502 @@ +package gigachat + +import ( + "context" + "errors" + "io" + "net/http" + "strings" + "sync" + "sync/atomic" + "time" + + openaiProvider "github.com/maximhq/bifrost/core/providers/openai" + providerUtils "github.com/maximhq/bifrost/core/providers/utils" + schemas "github.com/maximhq/bifrost/core/schemas" + "github.com/valyala/fasthttp" +) + +var _ schemas.Provider = (*GigaChatProvider)(nil) + +// GigaChatProvider implements the Provider interface for GigaChat's API. +type GigaChatProvider struct { + logger schemas.Logger + client *fasthttp.Client + streamingClient *fasthttp.Client + networkConfig schemas.NetworkConfig + sendBackRawRequest bool + sendBackRawResponse bool + customProviderConfig *schemas.CustomProviderConfig + tokenCache *gigaChatTokenCache + tlsClientCache *gigaChatTLSClientCache + attachmentCache *gigaChatAttachmentCacheManager +} + +// gigaChatPassthroughReadCloser keeps stream finalization aligned with the +// transport-owned large-response reader. The provider returns before the HTTP +// transport consumes that reader, so finalizing in the provider method would +// incorrectly mark a successful passthrough stream as incomplete. +type gigaChatPassthroughReadCloser struct { + io.ReadCloser + ctx *schemas.BifrostContext + postHookSpanFinalizer func(context.Context) + completed atomic.Bool + closeOnce sync.Once + closeErr error +} + +func (reader *gigaChatPassthroughReadCloser) Read(buffer []byte) (int, error) { + read, err := reader.ReadCloser.Read(buffer) + if errors.Is(err, io.EOF) { + reader.completed.Store(true) + } + return read, err +} + +func (reader *gigaChatPassthroughReadCloser) Close() error { + reader.closeOnce.Do(func() { + defer func() { + providerUtils.EnsureStreamFinalizerCalled(reader.ctx, reader.postHookSpanFinalizer) + }() + if reader.completed.Load() { + reader.ctx.SetValue(schemas.BifrostContextKeyStreamEndIndicator, true) + } + reader.closeErr = reader.ReadCloser.Close() + }) + return reader.closeErr +} + +func wrapGigaChatPassthroughFinalizer(ctx *schemas.BifrostContext, postHookSpanFinalizer func(context.Context)) { + reader, ok := ctx.Value(schemas.BifrostContextKeyLargeResponseReader).(io.ReadCloser) + if !ok || reader == nil { + providerUtils.EnsureStreamFinalizerCalled(ctx, postHookSpanFinalizer) + return + } + ctx.SetValue(schemas.BifrostContextKeyLargeResponseReader, &gigaChatPassthroughReadCloser{ + ReadCloser: reader, + ctx: ctx, + postHookSpanFinalizer: postHookSpanFinalizer, + }) +} + +// NewGigaChatProvider creates a new GigaChat provider instance. +func NewGigaChatProvider(config *schemas.ProviderConfig, logger schemas.Logger) (*GigaChatProvider, error) { + config.CheckAndSetDefaults() + + requestTimeout := time.Second * time.Duration(config.NetworkConfig.DefaultRequestTimeoutInSeconds) + client := &fasthttp.Client{ + ReadTimeout: requestTimeout, + WriteTimeout: requestTimeout, + MaxConnsPerHost: config.NetworkConfig.MaxConnsPerHost, + MaxIdleConnDuration: 30 * time.Second, + MaxConnWaitTimeout: requestTimeout, + MaxConnDuration: time.Second * time.Duration(schemas.DefaultMaxConnDurationInSeconds), + ConnPoolStrategy: fasthttp.FIFO, + } + + client = providerUtils.ConfigureProxy(client, config.ProxyConfig, logger) + client = providerUtils.ConfigureDialer(client, config.NetworkConfig.AllowPrivateNetwork) + client = providerUtils.ConfigureTLS(client, config.NetworkConfig, logger) + streamingClient := providerUtils.BuildStreamingClient(client) + + if config.NetworkConfig.BaseURL == "" { + config.NetworkConfig.BaseURL = gigaChatDefaultBaseURL + } + config.NetworkConfig.BaseURL = strings.TrimRight(config.NetworkConfig.BaseURL, "/") + + return &GigaChatProvider{ + logger: logger, + client: client, + streamingClient: streamingClient, + networkConfig: config.NetworkConfig, + sendBackRawRequest: config.SendBackRawRequest, + sendBackRawResponse: config.SendBackRawResponse, + customProviderConfig: config.CustomProviderConfig, + tokenCache: newGigaChatTokenCache(time.Now), + tlsClientCache: newGigaChatTLSClientCache(), + attachmentCache: newGigaChatAttachmentCacheManager(), + }, nil +} + +// GetProviderKey returns the provider identifier for GigaChat. +func (provider *GigaChatProvider) GetProviderKey() schemas.ModelProvider { + return providerUtils.GetProviderName(schemas.GigaChat, provider.customProviderConfig) +} + +func (provider *GigaChatProvider) chatCompletion(ctx *schemas.BifrostContext, key schemas.Key, request *schemas.BifrostChatRequest, forceRefresh bool) (*schemas.BifrostChatResponse, *schemas.BifrostError) { + ctx = ensureGigaChatContext(ctx) + if request == nil { + return nil, providerUtils.NewBifrostOperationError("chat completion request is nil", nil) + } + preparedRequest, bifrostErr := provider.prepareGigaChatChatAttachments(ctx, key, request) + if bifrostErr != nil { + return nil, bifrostErr + } + request = preparedRequest + + jsonData, bifrostErr := providerUtils.CheckContextAndGetRequestBody( + ctx, + request, + func() (providerUtils.RequestBodyWithExtraParams, error) { + return ToGigaChatChatRequest(ctx, request) + }) + if bifrostErr != nil { + return nil, bifrostErr + } + + headers, bifrostErr := provider.buildAuthHeadersWithRefresh(ctx, key, forceRefresh) + if bifrostErr != nil { + return nil, bifrostErr + } + + client, clientErr := provider.getGigaChatTLSClient(provider.client, gigaChatTLSClientCacheDefault, key.GigaChatKeyConfig) + if clientErr != nil { + return nil, newGigaChatConfigurationError(clientErr.Error()) + } + + req := fasthttp.AcquireRequest() + resp := fasthttp.AcquireResponse() + defer fasthttp.ReleaseRequest(req) + defer fasthttp.ReleaseResponse(resp) + + for headerName, headerValue := range headers { + req.Header.Set(headerName, headerValue) + } + req.SetRequestURI(buildGigaChatRequestURL(ctx, resolveBaseURL(key, provider.networkConfig), gigaChatAPIVersionV1, "/chat/completions", provider.customProviderConfig, schemas.ChatCompletionRequest)) + req.Header.SetMethod(http.MethodPost) + req.Header.SetContentType("application/json") + req.Header.Set("Accept", "application/json") + req.SetBody(jsonData) + + sendBackRawRequest := providerUtils.ShouldSendBackRawRequest(ctx, provider.sendBackRawRequest) + sendBackRawResponse := providerUtils.ShouldSendBackRawResponse(ctx, provider.sendBackRawResponse) + + latency, bifrostErr, wait := providerUtils.MakeRequestWithContext(ctx, client, req, resp) + defer wait() + if bifrostErr != nil { + bifrostErr.ExtraFields.Provider = provider.GetProviderKey() + return nil, enrichGigaChatError(ctx, bifrostErr, jsonData, nil, sendBackRawRequest, sendBackRawResponse) + } + + providerResponseHeaders := providerUtils.ExtractProviderResponseHeaders(resp) + ctx.SetValue(schemas.BifrostContextKeyProviderResponseHeaders, providerResponseHeaders) + + if resp.StatusCode() != fasthttp.StatusOK { + bifrostErr := ParseGigaChatError(resp, provider.GetProviderKey()) + return nil, enrichGigaChatError(ctx, bifrostErr, jsonData, nil, sendBackRawRequest, sendBackRawResponse) + } + + responseBody, err := providerUtils.CheckAndDecodeBody(resp) + if err != nil { + bifrostErr := newGigaChatProviderResponseError("failed to decode GigaChat chat completion response", err) + return nil, enrichGigaChatError(ctx, bifrostErr, jsonData, resp.Body(), sendBackRawRequest, sendBackRawResponse) + } + + gigaChatResponse := &GigaChatChatResponse{} + rawRequest, rawResponse, bifrostErr := providerUtils.HandleProviderResponse(responseBody, gigaChatResponse, jsonData, sendBackRawRequest, sendBackRawResponse) + if bifrostErr != nil { + return nil, enrichGigaChatError(ctx, bifrostErr, jsonData, responseBody, sendBackRawRequest, sendBackRawResponse) + } + + response := ToBifrostChatResponse(provider.GetProviderKey(), gigaChatResponse) + if response == nil { + return nil, newGigaChatProviderResponseError("GigaChat chat completion response is empty", nil) + } + response.BackfillParams(request) + response.ExtraFields.Latency = latency.Milliseconds() + response.ExtraFields.ProviderResponseHeaders = providerResponseHeaders + if sendBackRawRequest { + response.ExtraFields.RawRequest = rawRequest + } + if sendBackRawResponse { + response.ExtraFields.RawResponse = rawResponse + } + + return response, nil +} + +func (provider *GigaChatProvider) chatCompletionStream( + ctx *schemas.BifrostContext, + postHookRunner schemas.PostHookRunner, + postHookSpanFinalizer func(context.Context), + key schemas.Key, + request *schemas.BifrostChatRequest, + forceRefresh bool, +) (chan *schemas.BifrostStreamChunk, *schemas.BifrostError) { + ctx = ensureGigaChatContext(ctx) + if request == nil { + return nil, providerUtils.NewBifrostOperationError("chat completion request is nil", nil) + } + preparedRequest, bifrostErr := provider.prepareGigaChatChatAttachments(ctx, key, request) + if bifrostErr != nil { + return nil, bifrostErr + } + request = preparedRequest + + headers, bifrostErr := provider.buildAuthHeadersWithRefresh(ctx, key, forceRefresh) + if bifrostErr != nil { + return nil, bifrostErr + } + + client, clientErr := provider.getGigaChatTLSClient(provider.streamingClient, gigaChatTLSClientCacheStreaming, key.GigaChatKeyConfig) + if clientErr != nil { + return nil, newGigaChatConfigurationError(clientErr.Error()) + } + + sendBackRawRequest := providerUtils.ShouldSendBackRawRequest(ctx, provider.sendBackRawRequest) + sendBackRawResponse := providerUtils.ShouldSendBackRawResponse(ctx, provider.sendBackRawResponse) + providerName := provider.GetProviderKey() + + responseChan, bifrostErr := openaiProvider.HandleOpenAIChatCompletionStreaming( + ctx, + client, + buildGigaChatRequestURL(ctx, resolveBaseURL(key, provider.networkConfig), gigaChatAPIVersionV1, "/chat/completions", provider.customProviderConfig, schemas.ChatCompletionStreamRequest), + request, + headers, + nil, + provider.networkConfig.StreamIdleTimeoutInSeconds, + sendBackRawRequest, + sendBackRawResponse, + providerName, + postHookRunner, + func(request *schemas.BifrostChatRequest) (providerUtils.RequestBodyWithExtraParams, error) { + return ToGigaChatChatStreamRequest(ctx, request) + }, + handleGigaChatChatStreamResponse(providerName), + func(resp *fasthttp.Response) *schemas.BifrostError { + return ParseGigaChatError(resp, providerName) + }, + nil, + withGigaChatChatResponseProvider(providerName), + nil, + provider.logger, + postHookSpanFinalizer, + ) + if bifrostErr == nil { + if isPassthrough, _ := ctx.Value(schemas.BifrostContextKeyLargeResponseMode).(bool); isPassthrough { + wrapGigaChatPassthroughFinalizer(ctx, postHookSpanFinalizer) + } + } + return responseChan, bifrostErr +} + +func ensureGigaChatContext(ctx *schemas.BifrostContext) *schemas.BifrostContext { + if ctx != nil { + return ctx + } + return schemas.NewBifrostContext(context.Background(), schemas.NoDeadline) +} + +func isGigaChatUnauthorizedError(bifrostErr *schemas.BifrostError) bool { + return bifrostErr != nil && bifrostErr.StatusCode != nil && *bifrostErr.StatusCode == http.StatusUnauthorized +} + +func (provider *GigaChatProvider) unsupported(requestType schemas.RequestType) *schemas.BifrostError { + return providerUtils.NewUnsupportedOperationError(requestType, provider.GetProviderKey()) +} + +func (provider *GigaChatProvider) listModelsByKey(ctx *schemas.BifrostContext, key schemas.Key, request *schemas.BifrostListModelsRequest) (*schemas.BifrostListModelsResponse, *schemas.BifrostError) { + response, bifrostErr := provider.listModelsByKeyWithRefresh(ctx, key, request, false) + if isGigaChatUnauthorizedError(bifrostErr) { + return provider.listModelsByKeyWithRefresh(ctx, key, request, true) + } + return response, bifrostErr +} + +func (provider *GigaChatProvider) listModelsByKeyWithRefresh(ctx *schemas.BifrostContext, key schemas.Key, request *schemas.BifrostListModelsRequest, forceRefresh bool) (*schemas.BifrostListModelsResponse, *schemas.BifrostError) { + ctx = ensureGigaChatContext(ctx) + + headers, bifrostErr := provider.buildAuthHeadersWithRefresh(ctx, key, forceRefresh) + if bifrostErr != nil { + return nil, bifrostErr + } + + client, clientErr := provider.getGigaChatTLSClient(provider.client, gigaChatTLSClientCacheDefault, key.GigaChatKeyConfig) + if clientErr != nil { + return nil, newGigaChatConfigurationError(clientErr.Error()) + } + + req := fasthttp.AcquireRequest() + resp := fasthttp.AcquireResponse() + defer fasthttp.ReleaseRequest(req) + defer fasthttp.ReleaseResponse(resp) + + for headerName, headerValue := range headers { + req.Header.Set(headerName, headerValue) + } + req.SetRequestURI(buildGigaChatRequestURL(ctx, resolveBaseURL(key, provider.networkConfig), gigaChatAPIVersionV1, "/models", provider.customProviderConfig, schemas.ListModelsRequest)) + req.Header.SetMethod(http.MethodGet) + req.Header.SetContentType("application/json") + req.Header.Set("Accept", "application/json") + + sendBackRawRequest := providerUtils.ShouldSendBackRawRequest(ctx, provider.sendBackRawRequest) + sendBackRawResponse := providerUtils.ShouldSendBackRawResponse(ctx, provider.sendBackRawResponse) + + latency, bifrostErr, wait := providerUtils.MakeRequestWithContext(ctx, client, req, resp) + defer wait() + if bifrostErr != nil { + bifrostErr.ExtraFields.Provider = provider.GetProviderKey() + return nil, enrichGigaChatError(ctx, bifrostErr, nil, nil, sendBackRawRequest, sendBackRawResponse) + } + + providerResponseHeaders := providerUtils.ExtractProviderResponseHeaders(resp) + ctx.SetValue(schemas.BifrostContextKeyProviderResponseHeaders, providerResponseHeaders) + + if resp.StatusCode() != fasthttp.StatusOK { + bifrostErr := ParseGigaChatError(resp, provider.GetProviderKey()) + return nil, enrichGigaChatError(ctx, bifrostErr, nil, resp.Body(), sendBackRawRequest, sendBackRawResponse) + } + + responseBody, err := providerUtils.CheckAndDecodeBody(resp) + if err != nil { + bifrostErr := newGigaChatProviderResponseError("failed to decode GigaChat models response", err) + return nil, enrichGigaChatError(ctx, bifrostErr, nil, resp.Body(), sendBackRawRequest, sendBackRawResponse) + } + + gigaChatResponse := &GigaChatListModelsResponse{} + rawRequest, rawResponse, bifrostErr := providerUtils.HandleProviderResponse(responseBody, gigaChatResponse, nil, sendBackRawRequest, sendBackRawResponse) + if bifrostErr != nil { + return nil, enrichGigaChatError(ctx, bifrostErr, nil, responseBody, sendBackRawRequest, sendBackRawResponse) + } + + response := gigaChatResponse.ToBifrostListModelsResponse(provider.GetProviderKey(), key.Models, key.BlacklistedModels, key.Aliases, request.Unfiltered) + if response == nil { + return nil, newGigaChatProviderResponseError("GigaChat models response is empty", nil) + } + response.ExtraFields.Latency = latency.Milliseconds() + response.ExtraFields.ProviderResponseHeaders = providerResponseHeaders + if sendBackRawRequest { + response.ExtraFields.RawRequest = rawRequest + } + if sendBackRawResponse { + response.ExtraFields.RawResponse = rawResponse + } + + return response, nil +} + +func (provider *GigaChatProvider) embeddingWithRefresh(ctx *schemas.BifrostContext, key schemas.Key, request *schemas.BifrostEmbeddingRequest, forceRefresh bool) (*schemas.BifrostEmbeddingResponse, *schemas.BifrostError) { + ctx = ensureGigaChatContext(ctx) + if request == nil { + return nil, providerUtils.NewBifrostOperationError("embedding request is nil", nil) + } + + jsonData, bifrostErr := providerUtils.CheckContextAndGetRequestBody( + ctx, + request, + func() (providerUtils.RequestBodyWithExtraParams, error) { + return ToGigaChatEmbeddingRequest(request) + }) + if bifrostErr != nil { + return nil, bifrostErr + } + + headers, bifrostErr := provider.buildAuthHeadersWithRefresh(ctx, key, forceRefresh) + if bifrostErr != nil { + return nil, bifrostErr + } + + client, clientErr := provider.getGigaChatTLSClient(provider.client, gigaChatTLSClientCacheDefault, key.GigaChatKeyConfig) + if clientErr != nil { + return nil, newGigaChatConfigurationError(clientErr.Error()) + } + + req := fasthttp.AcquireRequest() + resp := fasthttp.AcquireResponse() + defer fasthttp.ReleaseRequest(req) + defer fasthttp.ReleaseResponse(resp) + + for headerName, headerValue := range headers { + req.Header.Set(headerName, headerValue) + } + req.SetRequestURI(buildGigaChatRequestURL(ctx, resolveBaseURL(key, provider.networkConfig), gigaChatAPIVersionV1, "/embeddings", provider.customProviderConfig, schemas.EmbeddingRequest)) + req.Header.SetMethod(http.MethodPost) + req.Header.SetContentType("application/json") + req.Header.Set("Accept", "application/json") + req.SetBody(jsonData) + + sendBackRawRequest := providerUtils.ShouldSendBackRawRequest(ctx, provider.sendBackRawRequest) + sendBackRawResponse := providerUtils.ShouldSendBackRawResponse(ctx, provider.sendBackRawResponse) + + latency, bifrostErr, wait := providerUtils.MakeRequestWithContext(ctx, client, req, resp) + defer wait() + if bifrostErr != nil { + bifrostErr.ExtraFields.Provider = provider.GetProviderKey() + return nil, enrichGigaChatError(ctx, bifrostErr, jsonData, nil, sendBackRawRequest, sendBackRawResponse) + } + + providerResponseHeaders := providerUtils.ExtractProviderResponseHeaders(resp) + ctx.SetValue(schemas.BifrostContextKeyProviderResponseHeaders, providerResponseHeaders) + + if resp.StatusCode() != fasthttp.StatusOK { + bifrostErr := ParseGigaChatError(resp, provider.GetProviderKey()) + return nil, enrichGigaChatError(ctx, bifrostErr, jsonData, resp.Body(), sendBackRawRequest, sendBackRawResponse) + } + + responseBody, err := providerUtils.CheckAndDecodeBody(resp) + if err != nil { + bifrostErr := newGigaChatProviderResponseError("failed to decode GigaChat embeddings response", err) + return nil, enrichGigaChatError(ctx, bifrostErr, jsonData, resp.Body(), sendBackRawRequest, sendBackRawResponse) + } + + gigaChatResponse := &GigaChatEmbeddingResponse{} + rawRequest, rawResponse, bifrostErr := providerUtils.HandleProviderResponse(responseBody, gigaChatResponse, jsonData, sendBackRawRequest, sendBackRawResponse) + if bifrostErr != nil { + return nil, enrichGigaChatError(ctx, bifrostErr, jsonData, responseBody, sendBackRawRequest, sendBackRawResponse) + } + + response := ToBifrostEmbeddingResponse(provider.GetProviderKey(), gigaChatResponse) + if response == nil { + return nil, newGigaChatProviderResponseError("GigaChat embeddings response is empty", nil) + } + if err := applyGigaChatEmbeddingEncodingFormat(response, request.Params); err != nil { + return nil, enrichGigaChatError(ctx, newGigaChatProviderResponseError("failed to encode GigaChat embeddings response", err), jsonData, responseBody, sendBackRawRequest, sendBackRawResponse) + } + response.BackfillParams(request) + response.ExtraFields.Latency = latency.Milliseconds() + response.ExtraFields.ProviderResponseHeaders = providerResponseHeaders + if sendBackRawRequest { + response.ExtraFields.RawRequest = rawRequest + } + if sendBackRawResponse { + response.ExtraFields.RawResponse = rawResponse + } + + return response, nil +} + +func (provider *GigaChatProvider) responsesWithRefresh(ctx *schemas.BifrostContext, key schemas.Key, request *schemas.BifrostResponsesRequest, forceRefresh bool) (*schemas.BifrostResponsesResponse, *schemas.BifrostError) { + ctx = ensureGigaChatContext(ctx) + if request == nil { + return nil, providerUtils.NewBifrostOperationError("responses request is nil", nil) + } + preparedRequest, bifrostErr := provider.prepareGigaChatResponsesAttachments(ctx, key, request) + if bifrostErr != nil { + return nil, bifrostErr + } + request = preparedRequest + + jsonData, bifrostErr := providerUtils.CheckContextAndGetRequestBody( + ctx, + request, + func() (providerUtils.RequestBodyWithExtraParams, error) { + return ToGigaChatResponsesRequest(request) + }) + if bifrostErr != nil { + return nil, bifrostErr + } + + headers, bifrostErr := provider.buildAuthHeadersWithRefresh(ctx, key, forceRefresh) + if bifrostErr != nil { + return nil, bifrostErr + } + + client, clientErr := provider.getGigaChatTLSClient(provider.client, gigaChatTLSClientCacheDefault, key.GigaChatKeyConfig) + if clientErr != nil { + return nil, newGigaChatConfigurationError(clientErr.Error()) + } + + req := fasthttp.AcquireRequest() + resp := fasthttp.AcquireResponse() + defer fasthttp.ReleaseRequest(req) + defer fasthttp.ReleaseResponse(resp) + + for headerName, headerValue := range headers { + req.Header.Set(headerName, headerValue) + } + req.SetRequestURI(buildGigaChatRequestURL(ctx, resolveBaseURL(key, provider.networkConfig), gigaChatAPIVersionV2, "/chat/completions", provider.customProviderConfig, schemas.ResponsesRequest)) + req.Header.SetMethod(http.MethodPost) + req.Header.SetContentType("application/json") + req.Header.Set("Accept", "application/json") + req.SetBody(jsonData) + + sendBackRawRequest := providerUtils.ShouldSendBackRawRequest(ctx, provider.sendBackRawRequest) + sendBackRawResponse := providerUtils.ShouldSendBackRawResponse(ctx, provider.sendBackRawResponse) + + latency, bifrostErr, wait := providerUtils.MakeRequestWithContext(ctx, client, req, resp) + defer wait() + if bifrostErr != nil { + bifrostErr.ExtraFields.Provider = provider.GetProviderKey() + return nil, enrichGigaChatError(ctx, bifrostErr, jsonData, nil, sendBackRawRequest, sendBackRawResponse) + } + + providerResponseHeaders := providerUtils.ExtractProviderResponseHeaders(resp) + ctx.SetValue(schemas.BifrostContextKeyProviderResponseHeaders, providerResponseHeaders) + + if resp.StatusCode() != fasthttp.StatusOK { + bifrostErr := ParseGigaChatError(resp, provider.GetProviderKey()) + return nil, enrichGigaChatError(ctx, bifrostErr, jsonData, resp.Body(), sendBackRawRequest, sendBackRawResponse) + } + + responseBody, err := providerUtils.CheckAndDecodeBody(resp) + if err != nil { + bifrostErr := newGigaChatProviderResponseError("failed to decode GigaChat Responses response", err) + return nil, enrichGigaChatError(ctx, bifrostErr, jsonData, resp.Body(), sendBackRawRequest, sendBackRawResponse) + } + + gigaChatResponse := &GigaChatResponsesResponse{} + rawRequest, rawResponse, bifrostErr := providerUtils.HandleProviderResponse(responseBody, gigaChatResponse, jsonData, sendBackRawRequest, sendBackRawResponse) + if bifrostErr != nil { + return nil, enrichGigaChatError(ctx, bifrostErr, jsonData, responseBody, sendBackRawRequest, sendBackRawResponse) + } + + response := ToBifrostResponsesResponse(provider.GetProviderKey(), gigaChatResponse) + if response == nil { + return nil, newGigaChatProviderResponseError("GigaChat Responses response is empty", nil) + } + response.BackfillParams(request) + response.ExtraFields.Latency = latency.Milliseconds() + response.ExtraFields.ProviderResponseHeaders = providerResponseHeaders + if sendBackRawRequest { + response.ExtraFields.RawRequest = rawRequest + } + if sendBackRawResponse { + response.ExtraFields.RawResponse = rawResponse + } + + return response, nil +} + +func (provider *GigaChatProvider) countTokensWithRefresh(ctx *schemas.BifrostContext, key schemas.Key, request *schemas.BifrostResponsesRequest, forceRefresh bool) (*schemas.BifrostCountTokensResponse, *schemas.BifrostError) { + ctx = ensureGigaChatContext(ctx) + if request == nil { + return nil, providerUtils.NewBifrostOperationError("count tokens request is nil", nil) + } + + jsonData, bifrostErr := providerUtils.CheckContextAndGetRequestBody( + ctx, + request, + func() (providerUtils.RequestBodyWithExtraParams, error) { + return ToGigaChatCountTokensRequest(request) + }) + if bifrostErr != nil { + return nil, bifrostErr + } + + headers, bifrostErr := provider.buildAuthHeadersWithRefresh(ctx, key, forceRefresh) + if bifrostErr != nil { + return nil, bifrostErr + } + + client, clientErr := provider.getGigaChatTLSClient(provider.client, gigaChatTLSClientCacheDefault, key.GigaChatKeyConfig) + if clientErr != nil { + return nil, newGigaChatConfigurationError(clientErr.Error()) + } + + req := fasthttp.AcquireRequest() + resp := fasthttp.AcquireResponse() + defer fasthttp.ReleaseRequest(req) + defer fasthttp.ReleaseResponse(resp) + + for headerName, headerValue := range headers { + req.Header.Set(headerName, headerValue) + } + req.SetRequestURI(buildGigaChatRequestURL(ctx, resolveBaseURL(key, provider.networkConfig), gigaChatAPIVersionV1, "/tokens/count", provider.customProviderConfig, schemas.CountTokensRequest)) + req.Header.SetMethod(http.MethodPost) + req.Header.SetContentType("application/json") + req.Header.Set("Accept", "application/json") + req.SetBody(jsonData) + + sendBackRawRequest := providerUtils.ShouldSendBackRawRequest(ctx, provider.sendBackRawRequest) + sendBackRawResponse := providerUtils.ShouldSendBackRawResponse(ctx, provider.sendBackRawResponse) + + latency, bifrostErr, wait := providerUtils.MakeRequestWithContext(ctx, client, req, resp) + defer wait() + if bifrostErr != nil { + bifrostErr.ExtraFields.Provider = provider.GetProviderKey() + return nil, enrichGigaChatError(ctx, bifrostErr, jsonData, nil, sendBackRawRequest, sendBackRawResponse) + } + + providerResponseHeaders := providerUtils.ExtractProviderResponseHeaders(resp) + ctx.SetValue(schemas.BifrostContextKeyProviderResponseHeaders, providerResponseHeaders) + + if resp.StatusCode() != fasthttp.StatusOK { + bifrostErr := ParseGigaChatError(resp, provider.GetProviderKey()) + return nil, enrichGigaChatError(ctx, bifrostErr, jsonData, resp.Body(), sendBackRawRequest, sendBackRawResponse) + } + + responseBody, err := providerUtils.CheckAndDecodeBody(resp) + if err != nil { + bifrostErr := newGigaChatProviderResponseError("failed to decode GigaChat count tokens response", err) + return nil, enrichGigaChatError(ctx, bifrostErr, jsonData, resp.Body(), sendBackRawRequest, sendBackRawResponse) + } + + gigaChatResponse := &GigaChatCountTokensResponse{} + rawRequest, rawResponse, bifrostErr := providerUtils.HandleProviderResponse(responseBody, gigaChatResponse, jsonData, sendBackRawRequest, sendBackRawResponse) + if bifrostErr != nil { + return nil, enrichGigaChatError(ctx, bifrostErr, jsonData, responseBody, sendBackRawRequest, sendBackRawResponse) + } + + response := ToBifrostCountTokensResponse(provider.GetProviderKey(), gigaChatResponse, request.Model) + if response == nil { + return nil, newGigaChatProviderResponseError("GigaChat count tokens response is empty", nil) + } + response.ExtraFields.Latency = latency.Milliseconds() + response.ExtraFields.ProviderResponseHeaders = providerResponseHeaders + if sendBackRawRequest { + response.ExtraFields.RawRequest = rawRequest + } + if sendBackRawResponse { + response.ExtraFields.RawResponse = rawResponse + } + + return response, nil +} + +func (provider *GigaChatProvider) responsesStreamWithRefresh( + ctx *schemas.BifrostContext, + postHookRunner schemas.PostHookRunner, + postHookSpanFinalizer func(context.Context), + key schemas.Key, + request *schemas.BifrostResponsesRequest, + forceRefresh bool, +) (chan *schemas.BifrostStreamChunk, *schemas.BifrostError) { + ctx = ensureGigaChatContext(ctx) + if request == nil { + return nil, providerUtils.NewBifrostOperationError("responses stream request is nil", nil) + } + preparedRequest, bifrostErr := provider.prepareGigaChatResponsesAttachments(ctx, key, request) + if bifrostErr != nil { + return nil, bifrostErr + } + request = preparedRequest + providerUtils.SetStreamIdleTimeoutIfEmpty(ctx, provider.networkConfig.StreamIdleTimeoutInSeconds) + + jsonData, bifrostErr := providerUtils.CheckContextAndGetRequestBody( + ctx, + request, + func() (providerUtils.RequestBodyWithExtraParams, error) { + return ToGigaChatResponsesStreamRequest(request) + }) + if bifrostErr != nil { + return nil, bifrostErr + } + + headers, bifrostErr := provider.buildAuthHeadersWithRefresh(ctx, key, forceRefresh) + if bifrostErr != nil { + return nil, bifrostErr + } + + client, clientErr := provider.getGigaChatTLSClient(provider.streamingClient, gigaChatTLSClientCacheStreaming, key.GigaChatKeyConfig) + if clientErr != nil { + return nil, newGigaChatConfigurationError(clientErr.Error()) + } + + req := fasthttp.AcquireRequest() + resp := fasthttp.AcquireResponse() + resp.StreamBody = true + defer fasthttp.ReleaseRequest(req) + + for headerName, headerValue := range headers { + req.Header.Set(headerName, headerValue) + } + req.SetRequestURI(buildGigaChatRequestURL(ctx, resolveBaseURL(key, provider.networkConfig), gigaChatAPIVersionV2, "/chat/completions", provider.customProviderConfig, schemas.ResponsesStreamRequest)) + req.Header.SetMethod(http.MethodPost) + req.Header.SetContentType("application/json") + req.Header.Set("Accept", "text/event-stream") + req.Header.Set("Cache-Control", "no-cache") + req.SetBody(jsonData) + + sendBackRawRequest := providerUtils.ShouldSendBackRawRequest(ctx, provider.sendBackRawRequest) + sendBackRawResponse := providerUtils.ShouldSendBackRawResponse(ctx, provider.sendBackRawResponse) + + activeClient := providerUtils.PrepareResponseStreaming(ctx, client, resp) + if err := activeClient.Do(req, resp); err != nil { + defer providerUtils.ReleaseStreamingResponse(ctx, resp) + if errors.Is(err, context.Canceled) { + return nil, providerUtils.EnrichError(ctx, &schemas.BifrostError{ + IsBifrostError: false, + Error: &schemas.ErrorField{ + Type: schemas.Ptr(schemas.RequestCancelled), + Message: schemas.ErrRequestCancelled, + Error: err, + }, + }, redactGigaChatRawPayload(jsonData), nil, sendBackRawRequest, sendBackRawResponse) + } + if errors.Is(err, fasthttp.ErrTimeout) || errors.Is(err, context.DeadlineExceeded) { + return nil, enrichGigaChatError(ctx, providerUtils.NewBifrostTimeoutError(schemas.ErrProviderRequestTimedOut, err), jsonData, nil, sendBackRawRequest, sendBackRawResponse) + } + return nil, enrichGigaChatError(ctx, providerUtils.NewBifrostOperationError(schemas.ErrProviderDoRequest, err), jsonData, nil, sendBackRawRequest, sendBackRawResponse) + } + startTime := time.Now() + + providerName := provider.GetProviderKey() + providerResponseHeaders := providerUtils.ExtractProviderResponseHeaders(resp) + ctx.SetValue(schemas.BifrostContextKeyProviderResponseHeaders, providerResponseHeaders) + + if resp.StatusCode() != fasthttp.StatusOK { + defer providerUtils.ReleaseStreamingResponse(ctx, resp) + providerUtils.MaterializeStreamErrorBody(ctx, resp) + bifrostErr := ParseGigaChatError(resp, providerName) + return nil, enrichGigaChatError(ctx, bifrostErr, jsonData, nil, sendBackRawRequest, sendBackRawResponse) + } + + if providerUtils.SetupStreamingPassthrough(ctx, resp) { + responseChan := make(chan *schemas.BifrostStreamChunk) + wrapGigaChatPassthroughFinalizer(ctx, postHookSpanFinalizer) + providerUtils.CloseStream(ctx, responseChan) + return responseChan, nil + } + + responseChan := make(chan *schemas.BifrostStreamChunk, schemas.DefaultStreamBufferSize) + + go func() { + defer providerUtils.EnsureStreamFinalizerCalled(ctx, postHookSpanFinalizer) + defer func() { + if ctx.Err() == context.Canceled { + providerUtils.HandleStreamCancellation(ctx, postHookRunner, responseChan, provider.logger, postHookSpanFinalizer, jsonData) + } else if ctx.Err() == context.DeadlineExceeded { + providerUtils.HandleStreamTimeout(ctx, postHookRunner, responseChan, provider.logger, postHookSpanFinalizer, jsonData) + } + providerUtils.CloseStream(ctx, responseChan) + }() + defer providerUtils.ReleaseStreamingResponse(ctx, resp) + + reader, releaseGzip := providerUtils.DecompressStreamBody(resp) + defer releaseGzip() + + reader, stopIdleTimeout := providerUtils.NewIdleTimeoutReader(reader, resp.BodyStream(), providerUtils.GetStreamIdleTimeout(ctx), ctx) + defer stopIdleTimeout() + + stopCancellation := providerUtils.SetupStreamCancellation(ctx, resp.BodyStream(), provider.logger) + defer stopCancellation() + + reader, drained := providerUtils.DrainNonSSEStreamReader(resp, reader) + if drained { + ctx.SetValue(schemas.BifrostContextKeyStreamEndIndicator, true) + providerUtils.ProcessAndSendError(ctx, postHookRunner, errors.New("provider returned non-SSE response for streaming request"), responseChan, provider.logger, postHookSpanFinalizer) + return + } + + sseReader := providerUtils.GetSSEEventReader(ctx, reader) + streamState := schemas.AcquireChatToResponsesStreamState() + defer schemas.ReleaseChatToResponsesStreamState(streamState) + + usage := &schemas.BifrostLLMUsage{} + usageSeen := false + lastChunkTime := startTime + var pendingFinalEvent *schemas.BifrostResponsesStreamResponse + streamEndedSemantically := false + + for { + if ctx.Err() != nil { + return + } + + eventType, data, readErr := sseReader.ReadEvent() + if readErr != nil { + if ctx.Err() != nil { + return + } + if readErr != io.EOF { + ctx.SetValue(schemas.BifrostContextKeyStreamEndIndicator, true) + if provider.logger != nil { + provider.logger.Warn("Error reading stream: %v", readErr) + } + providerUtils.ProcessAndSendError(ctx, postHookRunner, readErr, responseChan, provider.logger, postHookSpanFinalizer) + return + } + break + } + if isGigaChatResponsesStreamDoneMarker(data) { + streamEndedSemantically = true + break + } + if len(data) == 0 { + if isGigaChatResponsesStreamTerminalEvent(eventType, nil) { + streamEndedSemantically = true + break + } + continue + } + + if bifrostErr := parseGigaChatStreamError(data, providerName); bifrostErr != nil { + ctx.SetValue(schemas.BifrostContextKeyStreamEndIndicator, true) + providerUtils.ProcessAndSendBifrostError(ctx, postHookRunner, enrichGigaChatError(ctx, bifrostErr, jsonData, data, sendBackRawRequest, sendBackRawResponse), responseChan, provider.logger, postHookSpanFinalizer) + return + } + + var gigaChatResponse GigaChatResponsesResponse + _, rawResponse, handlerErr := providerUtils.HandleProviderResponse(data, &gigaChatResponse, nil, false, sendBackRawResponse) + if handlerErr != nil { + ctx.SetValue(schemas.BifrostContextKeyStreamEndIndicator, true) + providerUtils.ProcessAndSendBifrostError(ctx, postHookRunner, enrichGigaChatError(ctx, handlerErr, jsonData, data, sendBackRawRequest, sendBackRawResponse), responseChan, provider.logger, postHookSpanFinalizer) + return + } + + if gigaChatResponse.Usage != nil { + usageSeen = true + updateGigaChatResponsesStreamUsage(usage, toBifrostGigaChatUsage(gigaChatResponse.Usage)) + } + + responses := ToBifrostResponsesStreamResponse(providerName, &gigaChatResponse, streamState) + for _, response := range responses { + if response == nil { + continue + } + response.ExtraFields.ChunkIndex = response.SequenceNumber + response.ExtraFields.ProviderResponseHeaders = providerResponseHeaders + if sendBackRawResponse { + response.ExtraFields.RawResponse = rawResponse + } + + if response.Type == schemas.ResponsesStreamResponseTypeCompleted || response.Type == schemas.ResponsesStreamResponseTypeIncomplete { + pendingFinalEvent = response + continue + } + + response.ExtraFields.Latency = time.Since(lastChunkTime).Milliseconds() + lastChunkTime = time.Now() + providerUtils.ProcessAndSendResponse(ctx, postHookRunner, providerUtils.GetBifrostResponseForStreamResponse(nil, nil, response, nil, nil, nil), responseChan, postHookSpanFinalizer) + } + if isGigaChatResponsesStreamTerminalEvent(eventType, &gigaChatResponse) { + streamEndedSemantically = true + break + } + } + + if pendingFinalEvent != nil { + if usageSeen && pendingFinalEvent.Response != nil { + pendingFinalEvent.Response.Usage = usage.ToResponsesResponseUsage() + } + if sendBackRawRequest { + providerUtils.ParseAndSetRawRequest(&pendingFinalEvent.ExtraFields, jsonData) + } + pendingFinalEvent.ExtraFields.Latency = time.Since(startTime).Milliseconds() + ctx.SetValue(schemas.BifrostContextKeyStreamEndIndicator, true) + providerUtils.ProcessAndSendResponse(ctx, postHookRunner, providerUtils.GetBifrostResponseForStreamResponse(nil, nil, pendingFinalEvent, nil, nil, nil), responseChan, postHookSpanFinalizer) + } + if streamEndedSemantically { + closeGigaChatSemanticStream(ctx, resp.BodyStream()) + } + }() + + return responseChan, nil +} + +type gigaChatStreamCloserWithError interface { + CloseWithError(error) error +} + +func isGigaChatResponsesStreamDoneMarker(data []byte) bool { + return strings.TrimSpace(string(data)) == "[DONE]" +} + +func isGigaChatResponsesStreamTerminalEvent(eventType string, response *GigaChatResponsesResponse) bool { + if isGigaChatResponsesStreamTerminalEventName(eventType) { + return true + } + if response == nil { + return false + } + if response.Event != nil && isGigaChatResponsesStreamTerminalEventName(*response.Event) { + return true + } + return response.FinishReason != nil && len(response.Messages) == 0 && len(response.Choices) == 0 +} + +func isGigaChatResponsesStreamTerminalEventName(eventType string) bool { + switch strings.ToLower(strings.TrimSpace(eventType)) { + case "done", "response.done", "response.completed", "response.message.done": + return true + default: + return false + } +} + +func closeGigaChatSemanticStream(ctx *schemas.BifrostContext, bodyStream io.Reader) { + if bodyStream == nil { + return + } + if closed, ok := ctx.Value(schemas.BifrostContextKeyConnectionClosed).(bool); ok && closed { + return + } + if closer, ok := bodyStream.(io.Closer); ok { + ctx.SetValue(schemas.BifrostContextKeyConnectionClosed, true) + _ = closer.Close() + return + } + if closer, ok := bodyStream.(gigaChatStreamCloserWithError); ok { + ctx.SetValue(schemas.BifrostContextKeyConnectionClosed, true) + _ = closer.CloseWithError(io.EOF) + } +} + +// ListModels performs a v1 models request to GigaChat. +func (provider *GigaChatProvider) ListModels(ctx *schemas.BifrostContext, keys []schemas.Key, request *schemas.BifrostListModelsRequest) (*schemas.BifrostListModelsResponse, *schemas.BifrostError) { + if err := providerUtils.CheckOperationAllowed(schemas.GigaChat, provider.customProviderConfig, schemas.ListModelsRequest); err != nil { + return nil, err + } + if request == nil { + request = &schemas.BifrostListModelsRequest{Provider: provider.GetProviderKey()} + } else if request.Provider == "" { + requestCopy := *request + requestCopy.Provider = provider.GetProviderKey() + request = &requestCopy + } + if len(keys) == 0 { + return providerUtils.HandleKeylessListModelsRequest(provider.GetProviderKey(), func() (*schemas.BifrostListModelsResponse, *schemas.BifrostError) { + return provider.listModelsByKey(ctx, schemas.Key{}, request) + }) + } + return providerUtils.HandleMultipleListModelsRequests(ctx, keys, request, provider.listModelsByKey) +} + +// TextCompletion is not supported by the GigaChat provider skeleton. +func (provider *GigaChatProvider) TextCompletion(_ *schemas.BifrostContext, _ schemas.Key, _ *schemas.BifrostTextCompletionRequest) (*schemas.BifrostTextCompletionResponse, *schemas.BifrostError) { + return nil, provider.unsupported(schemas.TextCompletionRequest) +} + +// TextCompletionStream is not supported by the GigaChat provider skeleton. +func (provider *GigaChatProvider) TextCompletionStream(_ *schemas.BifrostContext, _ schemas.PostHookRunner, _ func(context.Context), _ schemas.Key, _ *schemas.BifrostTextCompletionRequest) (chan *schemas.BifrostStreamChunk, *schemas.BifrostError) { + return nil, provider.unsupported(schemas.TextCompletionStreamRequest) +} + +// ChatCompletion sends a non-streaming v1 chat completions request to GigaChat. +func (provider *GigaChatProvider) ChatCompletion(ctx *schemas.BifrostContext, key schemas.Key, request *schemas.BifrostChatRequest) (*schemas.BifrostChatResponse, *schemas.BifrostError) { + if err := providerUtils.CheckOperationAllowed(schemas.GigaChat, provider.customProviderConfig, schemas.ChatCompletionRequest); err != nil { + return nil, err + } + + response, bifrostErr := provider.chatCompletion(ctx, key, request, false) + if isGigaChatUnauthorizedError(bifrostErr) { + return provider.chatCompletion(ctx, key, request, true) + } + return response, bifrostErr +} + +// ChatCompletionStream sends a streaming v1 chat completions request to GigaChat. +func (provider *GigaChatProvider) ChatCompletionStream(ctx *schemas.BifrostContext, postHookRunner schemas.PostHookRunner, postHookSpanFinalizer func(context.Context), key schemas.Key, request *schemas.BifrostChatRequest) (chan *schemas.BifrostStreamChunk, *schemas.BifrostError) { + if err := providerUtils.CheckOperationAllowed(schemas.GigaChat, provider.customProviderConfig, schemas.ChatCompletionStreamRequest); err != nil { + return nil, err + } + + responseChan, bifrostErr := provider.chatCompletionStream(ctx, postHookRunner, postHookSpanFinalizer, key, request, false) + if isGigaChatUnauthorizedError(bifrostErr) { + return provider.chatCompletionStream(ctx, postHookRunner, postHookSpanFinalizer, key, request, true) + } + return responseChan, bifrostErr +} + +// Responses sends a non-streaming v2 chat completions request to GigaChat. +func (provider *GigaChatProvider) Responses(ctx *schemas.BifrostContext, key schemas.Key, request *schemas.BifrostResponsesRequest) (*schemas.BifrostResponsesResponse, *schemas.BifrostError) { + if err := providerUtils.CheckOperationAllowed(schemas.GigaChat, provider.customProviderConfig, schemas.ResponsesRequest); err != nil { + return nil, err + } + + response, bifrostErr := provider.responsesWithRefresh(ctx, key, request, false) + if isGigaChatUnauthorizedError(bifrostErr) { + return provider.responsesWithRefresh(ctx, key, request, true) + } + return response, bifrostErr +} + +// ResponsesStream sends a streaming v2 chat completions request to GigaChat. +func (provider *GigaChatProvider) ResponsesStream(ctx *schemas.BifrostContext, postHookRunner schemas.PostHookRunner, postHookSpanFinalizer func(context.Context), key schemas.Key, request *schemas.BifrostResponsesRequest) (chan *schemas.BifrostStreamChunk, *schemas.BifrostError) { + if err := providerUtils.CheckOperationAllowed(schemas.GigaChat, provider.customProviderConfig, schemas.ResponsesStreamRequest); err != nil { + return nil, err + } + + responseChan, bifrostErr := provider.responsesStreamWithRefresh(ctx, postHookRunner, postHookSpanFinalizer, key, request, false) + if isGigaChatUnauthorizedError(bifrostErr) { + return provider.responsesStreamWithRefresh(ctx, postHookRunner, postHookSpanFinalizer, key, request, true) + } + return responseChan, bifrostErr +} + +// CountTokens sends a v1 tokens/count request to GigaChat. +func (provider *GigaChatProvider) CountTokens(ctx *schemas.BifrostContext, key schemas.Key, request *schemas.BifrostResponsesRequest) (*schemas.BifrostCountTokensResponse, *schemas.BifrostError) { + if err := providerUtils.CheckOperationAllowed(schemas.GigaChat, provider.customProviderConfig, schemas.CountTokensRequest); err != nil { + return nil, err + } + + response, bifrostErr := provider.countTokensWithRefresh(ctx, key, request, false) + if isGigaChatUnauthorizedError(bifrostErr) { + return provider.countTokensWithRefresh(ctx, key, request, true) + } + return response, bifrostErr +} + +// Compaction is not supported by the GigaChat provider skeleton. +func (provider *GigaChatProvider) Compaction(_ *schemas.BifrostContext, _ schemas.Key, _ *schemas.BifrostCompactionRequest) (*schemas.BifrostCompactionResponse, *schemas.BifrostError) { + return nil, provider.unsupported(schemas.CompactionRequest) +} + +// Embedding sends a non-streaming v1 embeddings request to GigaChat. +func (provider *GigaChatProvider) Embedding(ctx *schemas.BifrostContext, key schemas.Key, request *schemas.BifrostEmbeddingRequest) (*schemas.BifrostEmbeddingResponse, *schemas.BifrostError) { + if err := providerUtils.CheckOperationAllowed(schemas.GigaChat, provider.customProviderConfig, schemas.EmbeddingRequest); err != nil { + return nil, err + } + + response, bifrostErr := provider.embeddingWithRefresh(ctx, key, request, false) + if isGigaChatUnauthorizedError(bifrostErr) { + return provider.embeddingWithRefresh(ctx, key, request, true) + } + return response, bifrostErr +} + +// Rerank is not supported by the GigaChat provider skeleton. +func (provider *GigaChatProvider) Rerank(_ *schemas.BifrostContext, _ schemas.Key, _ *schemas.BifrostRerankRequest) (*schemas.BifrostRerankResponse, *schemas.BifrostError) { + return nil, provider.unsupported(schemas.RerankRequest) +} + +// OCR is not supported by the GigaChat provider skeleton. +func (provider *GigaChatProvider) OCR(_ *schemas.BifrostContext, _ schemas.Key, _ *schemas.BifrostOCRRequest) (*schemas.BifrostOCRResponse, *schemas.BifrostError) { + return nil, provider.unsupported(schemas.OCRRequest) +} + +// Speech is not supported by the GigaChat provider skeleton. +func (provider *GigaChatProvider) Speech(_ *schemas.BifrostContext, _ schemas.Key, _ *schemas.BifrostSpeechRequest) (*schemas.BifrostSpeechResponse, *schemas.BifrostError) { + return nil, provider.unsupported(schemas.SpeechRequest) +} + +// SpeechStream is not supported by the GigaChat provider skeleton. +func (provider *GigaChatProvider) SpeechStream(_ *schemas.BifrostContext, _ schemas.PostHookRunner, _ func(context.Context), _ schemas.Key, _ *schemas.BifrostSpeechRequest) (chan *schemas.BifrostStreamChunk, *schemas.BifrostError) { + return nil, provider.unsupported(schemas.SpeechStreamRequest) +} + +// Transcription is not supported by the GigaChat provider skeleton. +func (provider *GigaChatProvider) Transcription(_ *schemas.BifrostContext, _ schemas.Key, _ *schemas.BifrostTranscriptionRequest) (*schemas.BifrostTranscriptionResponse, *schemas.BifrostError) { + return nil, provider.unsupported(schemas.TranscriptionRequest) +} + +// TranscriptionStream is not supported by the GigaChat provider skeleton. +func (provider *GigaChatProvider) TranscriptionStream(_ *schemas.BifrostContext, _ schemas.PostHookRunner, _ func(context.Context), _ schemas.Key, _ *schemas.BifrostTranscriptionRequest) (chan *schemas.BifrostStreamChunk, *schemas.BifrostError) { + return nil, provider.unsupported(schemas.TranscriptionStreamRequest) +} + +// ImageGeneration is not supported by the GigaChat provider skeleton. +func (provider *GigaChatProvider) ImageGeneration(_ *schemas.BifrostContext, _ schemas.Key, _ *schemas.BifrostImageGenerationRequest) (*schemas.BifrostImageGenerationResponse, *schemas.BifrostError) { + return nil, provider.unsupported(schemas.ImageGenerationRequest) +} + +// ImageGenerationStream is not supported by the GigaChat provider skeleton. +func (provider *GigaChatProvider) ImageGenerationStream(_ *schemas.BifrostContext, _ schemas.PostHookRunner, _ func(context.Context), _ schemas.Key, _ *schemas.BifrostImageGenerationRequest) (chan *schemas.BifrostStreamChunk, *schemas.BifrostError) { + return nil, provider.unsupported(schemas.ImageGenerationStreamRequest) +} + +// ImageEdit is not supported by the GigaChat provider skeleton. +func (provider *GigaChatProvider) ImageEdit(_ *schemas.BifrostContext, _ schemas.Key, _ *schemas.BifrostImageEditRequest) (*schemas.BifrostImageGenerationResponse, *schemas.BifrostError) { + return nil, provider.unsupported(schemas.ImageEditRequest) +} + +// ImageEditStream is not supported by the GigaChat provider skeleton. +func (provider *GigaChatProvider) ImageEditStream(_ *schemas.BifrostContext, _ schemas.PostHookRunner, _ func(context.Context), _ schemas.Key, _ *schemas.BifrostImageEditRequest) (chan *schemas.BifrostStreamChunk, *schemas.BifrostError) { + return nil, provider.unsupported(schemas.ImageEditStreamRequest) +} + +// ImageVariation is not supported by the GigaChat provider skeleton. +func (provider *GigaChatProvider) ImageVariation(_ *schemas.BifrostContext, _ schemas.Key, _ *schemas.BifrostImageVariationRequest) (*schemas.BifrostImageGenerationResponse, *schemas.BifrostError) { + return nil, provider.unsupported(schemas.ImageVariationRequest) +} + +// VideoGeneration is not supported by the GigaChat provider skeleton. +func (provider *GigaChatProvider) VideoGeneration(_ *schemas.BifrostContext, _ schemas.Key, _ *schemas.BifrostVideoGenerationRequest) (*schemas.BifrostVideoGenerationResponse, *schemas.BifrostError) { + return nil, provider.unsupported(schemas.VideoGenerationRequest) +} + +// VideoEdit is not supported by the GigaChat provider skeleton. +func (provider *GigaChatProvider) VideoEdit(_ *schemas.BifrostContext, _ schemas.Key, _ *schemas.BifrostVideoEditRequest) (*schemas.BifrostVideoEditResponse, *schemas.BifrostError) { + return nil, provider.unsupported(schemas.VideoEditRequest) +} + +// VideoRetrieve is not supported by the GigaChat provider skeleton. +func (provider *GigaChatProvider) VideoRetrieve(_ *schemas.BifrostContext, _ schemas.Key, _ *schemas.BifrostVideoRetrieveRequest) (*schemas.BifrostVideoGenerationResponse, *schemas.BifrostError) { + return nil, provider.unsupported(schemas.VideoRetrieveRequest) +} + +// VideoDownload is not supported by the GigaChat provider skeleton. +func (provider *GigaChatProvider) VideoDownload(_ *schemas.BifrostContext, _ schemas.Key, _ *schemas.BifrostVideoDownloadRequest) (*schemas.BifrostVideoDownloadResponse, *schemas.BifrostError) { + return nil, provider.unsupported(schemas.VideoDownloadRequest) +} + +// VideoDelete is not supported by the GigaChat provider skeleton. +func (provider *GigaChatProvider) VideoDelete(_ *schemas.BifrostContext, _ schemas.Key, _ *schemas.BifrostVideoDeleteRequest) (*schemas.BifrostVideoDeleteResponse, *schemas.BifrostError) { + return nil, provider.unsupported(schemas.VideoDeleteRequest) +} + +// VideoList is not supported by the GigaChat provider skeleton. +func (provider *GigaChatProvider) VideoList(_ *schemas.BifrostContext, _ schemas.Key, _ *schemas.BifrostVideoListRequest) (*schemas.BifrostVideoListResponse, *schemas.BifrostError) { + return nil, provider.unsupported(schemas.VideoListRequest) +} + +// VideoRemix is not supported by the GigaChat provider skeleton. +func (provider *GigaChatProvider) VideoRemix(_ *schemas.BifrostContext, _ schemas.Key, _ *schemas.BifrostVideoRemixRequest) (*schemas.BifrostVideoGenerationResponse, *schemas.BifrostError) { + return nil, provider.unsupported(schemas.VideoRemixRequest) +} + +// BatchCreate creates a GigaChat batch job. +func (provider *GigaChatProvider) BatchCreate(ctx *schemas.BifrostContext, key schemas.Key, request *schemas.BifrostBatchCreateRequest) (*schemas.BifrostBatchCreateResponse, *schemas.BifrostError) { + if err := providerUtils.CheckOperationAllowed(schemas.GigaChat, provider.customProviderConfig, schemas.BatchCreateRequest); err != nil { + return nil, err + } + + response, bifrostErr := provider.batchCreateWithRefresh(ctx, key, request, false) + if isGigaChatUnauthorizedError(bifrostErr) { + return provider.batchCreateWithRefresh(ctx, key, request, true) + } + return response, bifrostErr +} + +// BatchList lists GigaChat batch jobs. +func (provider *GigaChatProvider) BatchList(ctx *schemas.BifrostContext, keys []schemas.Key, request *schemas.BifrostBatchListRequest) (*schemas.BifrostBatchListResponse, *schemas.BifrostError) { + if err := providerUtils.CheckOperationAllowed(schemas.GigaChat, provider.customProviderConfig, schemas.BatchListRequest); err != nil { + return nil, err + } + if request == nil { + request = &schemas.BifrostBatchListRequest{Provider: provider.GetProviderKey()} + } + if bifrostErr := validateGigaChatBatchListRequest(request); bifrostErr != nil { + return nil, bifrostErr + } + if len(keys) == 0 { + keys = []schemas.Key{{}} + } + + helper, err := providerUtils.NewSerialListHelper(keys, request.After, provider.logger, false) + if err != nil { + return nil, providerUtils.NewBifrostOperationError("invalid pagination cursor", err) + } + key, nativeCursor, ok := helper.GetCurrentKey() + if !ok { + return &schemas.BifrostBatchListResponse{ + Object: "list", + Data: []schemas.BifrostBatchRetrieveResponse{}, + }, nil + } + + response, bifrostErr := provider.batchListWithRefresh(ctx, key, request, false) + if isGigaChatUnauthorizedError(bifrostErr) { + response, bifrostErr = provider.batchListWithRefresh(ctx, key, request, true) + } + if bifrostErr != nil { + return nil, bifrostErr + } + + nativeNextCursor, nativeHasMore, err := paginateGigaChatBatchList(response, nativeCursor, request.Limit) + if err != nil { + return nil, providerUtils.NewBifrostOperationError("invalid pagination cursor", err) + } + nextCursor, hasMore := helper.BuildNextCursor(nativeHasMore, nativeNextCursor) + response.HasMore = hasMore + if nextCursor != "" { + response.NextCursor = &nextCursor + } + return response, nil +} + +// BatchRetrieve retrieves a GigaChat batch job. +func (provider *GigaChatProvider) BatchRetrieve(ctx *schemas.BifrostContext, keys []schemas.Key, request *schemas.BifrostBatchRetrieveRequest) (*schemas.BifrostBatchRetrieveResponse, *schemas.BifrostError) { + if err := providerUtils.CheckOperationAllowed(schemas.GigaChat, provider.customProviderConfig, schemas.BatchRetrieveRequest); err != nil { + return nil, err + } + if request == nil { + return nil, providerUtils.NewBifrostOperationError("batch retrieve request is nil", nil) + } + if strings.TrimSpace(request.BatchID) == "" { + return nil, providerUtils.NewBifrostOperationError("batch_id is required", nil) + } + if len(keys) == 0 { + keys = []schemas.Key{{}} + } + + var lastErr *schemas.BifrostError + for _, key := range keys { + response, bifrostErr := provider.batchRetrieveWithRefresh(ctx, key, request, false) + if isGigaChatUnauthorizedError(bifrostErr) { + response, bifrostErr = provider.batchRetrieveWithRefresh(ctx, key, request, true) + } + if bifrostErr == nil { + return response, nil + } + lastErr = bifrostErr + } + return nil, lastErr +} + +// BatchCancel is not supported by the GigaChat provider skeleton. +func (provider *GigaChatProvider) BatchCancel(_ *schemas.BifrostContext, _ []schemas.Key, _ *schemas.BifrostBatchCancelRequest) (*schemas.BifrostBatchCancelResponse, *schemas.BifrostError) { + return nil, provider.unsupported(schemas.BatchCancelRequest) +} + +// BatchDelete is not supported by the GigaChat provider skeleton. +func (provider *GigaChatProvider) BatchDelete(_ *schemas.BifrostContext, _ []schemas.Key, _ *schemas.BifrostBatchDeleteRequest) (*schemas.BifrostBatchDeleteResponse, *schemas.BifrostError) { + return nil, provider.unsupported(schemas.BatchDeleteRequest) +} + +// BatchResults retrieves completed GigaChat batch results through the Files API. +func (provider *GigaChatProvider) BatchResults(ctx *schemas.BifrostContext, keys []schemas.Key, request *schemas.BifrostBatchResultsRequest) (*schemas.BifrostBatchResultsResponse, *schemas.BifrostError) { + if err := providerUtils.CheckOperationAllowed(schemas.GigaChat, provider.customProviderConfig, schemas.BatchResultsRequest); err != nil { + return nil, err + } + if request == nil { + return nil, providerUtils.NewBifrostOperationError("batch results request is nil", nil) + } + if strings.TrimSpace(request.BatchID) == "" { + return nil, providerUtils.NewBifrostOperationError("batch_id is required", nil) + } + if len(keys) == 0 { + keys = []schemas.Key{{}} + } + + batchResponse, bifrostErr := provider.BatchRetrieve(ctx, keys, &schemas.BifrostBatchRetrieveRequest{ + Provider: request.Provider, + Model: request.Model, + BatchID: strings.TrimSpace(request.BatchID), + }) + if bifrostErr != nil { + return nil, bifrostErr + } + if batchResponse.OutputFileID == nil || strings.TrimSpace(*batchResponse.OutputFileID) == "" { + return nil, providerUtils.NewBifrostOperationError("batch results not available: GigaChat did not return output_file_id or result_file_id (batch may not be completed yet)", nil) + } + + outputFileID := strings.TrimSpace(*batchResponse.OutputFileID) + fileContentResponse, bifrostErr := provider.readGigaChatBatchOutputFile(ctx, keys, request, outputFileID) + if bifrostErr != nil { + return nil, bifrostErr + } + + results, parseErrors := parseGigaChatBatchResultsJSONL(fileContentResponse.Content, provider.logger) + + response := &schemas.BifrostBatchResultsResponse{ + BatchID: strings.TrimSpace(request.BatchID), + Results: results, + ExtraFields: schemas.BifrostResponseExtraFields{ + Provider: provider.GetProviderKey(), + Latency: fileContentResponse.ExtraFields.Latency, + ProviderResponseHeaders: fileContentResponse.ExtraFields.ProviderResponseHeaders, + }, + } + if len(parseErrors) > 0 { + response.ExtraFields.ParseErrors = parseErrors + } + return response, nil +} + +// FileUpload uploads a file to GigaChat. +func (provider *GigaChatProvider) FileUpload(ctx *schemas.BifrostContext, key schemas.Key, request *schemas.BifrostFileUploadRequest) (*schemas.BifrostFileUploadResponse, *schemas.BifrostError) { + if err := providerUtils.CheckOperationAllowed(schemas.GigaChat, provider.customProviderConfig, schemas.FileUploadRequest); err != nil { + return nil, err + } + + response, bifrostErr := provider.fileUploadWithRefresh(ctx, key, request, false) + if isGigaChatUnauthorizedError(bifrostErr) { + return provider.fileUploadWithRefresh(ctx, key, request, true) + } + return response, bifrostErr +} + +// FileList lists files available to the configured GigaChat account. +func (provider *GigaChatProvider) FileList(ctx *schemas.BifrostContext, keys []schemas.Key, request *schemas.BifrostFileListRequest) (*schemas.BifrostFileListResponse, *schemas.BifrostError) { + if err := providerUtils.CheckOperationAllowed(schemas.GigaChat, provider.customProviderConfig, schemas.FileListRequest); err != nil { + return nil, err + } + if request == nil { + request = &schemas.BifrostFileListRequest{Provider: provider.GetProviderKey()} + } + if request.Order != nil && strings.TrimSpace(*request.Order) != "" { + return nil, providerUtils.NewBifrostOperationError("GigaChat file list does not support order sorting", nil) + } + if len(keys) == 0 { + keys = []schemas.Key{{}} + } + + helper, err := providerUtils.NewSerialListHelper(keys, request.After, provider.logger, false) + if err != nil { + return nil, providerUtils.NewBifrostOperationError("invalid pagination cursor", err) + } + key, nativeCursor, ok := helper.GetCurrentKey() + if !ok { + return &schemas.BifrostFileListResponse{ + Object: "list", + Data: []schemas.FileObject{}, + }, nil + } + + response, bifrostErr := provider.fileListWithRefresh(ctx, key, request, false) + if isGigaChatUnauthorizedError(bifrostErr) { + response, bifrostErr = provider.fileListWithRefresh(ctx, key, request, true) + } + if bifrostErr != nil { + return nil, bifrostErr + } + + nativeNextCursor, nativeHasMore, err := paginateGigaChatFileList(response, nativeCursor, request.Limit) + if err != nil { + return nil, providerUtils.NewBifrostOperationError("invalid pagination cursor", err) + } + nextCursor, hasMore := helper.BuildNextCursor(nativeHasMore, nativeNextCursor) + response.HasMore = hasMore + if nextCursor != "" { + response.After = &nextCursor + } + return response, nil +} + +// FileRetrieve retrieves GigaChat file metadata. +func (provider *GigaChatProvider) FileRetrieve(ctx *schemas.BifrostContext, keys []schemas.Key, request *schemas.BifrostFileRetrieveRequest) (*schemas.BifrostFileRetrieveResponse, *schemas.BifrostError) { + if err := providerUtils.CheckOperationAllowed(schemas.GigaChat, provider.customProviderConfig, schemas.FileRetrieveRequest); err != nil { + return nil, err + } + if request == nil { + return nil, providerUtils.NewBifrostOperationError("file retrieve request is nil", nil) + } + if strings.TrimSpace(request.FileID) == "" { + return nil, providerUtils.NewBifrostOperationError("file_id is required", nil) + } + if len(keys) == 0 { + keys = []schemas.Key{{}} + } + + var lastErr *schemas.BifrostError + for _, key := range keys { + response, bifrostErr := provider.fileRetrieveWithRefresh(ctx, key, request, false) + if isGigaChatUnauthorizedError(bifrostErr) { + response, bifrostErr = provider.fileRetrieveWithRefresh(ctx, key, request, true) + } + if bifrostErr == nil { + return response, nil + } + lastErr = bifrostErr + } + return nil, lastErr +} + +// FileDelete deletes a GigaChat file. +func (provider *GigaChatProvider) FileDelete(ctx *schemas.BifrostContext, keys []schemas.Key, request *schemas.BifrostFileDeleteRequest) (*schemas.BifrostFileDeleteResponse, *schemas.BifrostError) { + if err := providerUtils.CheckOperationAllowed(schemas.GigaChat, provider.customProviderConfig, schemas.FileDeleteRequest); err != nil { + return nil, err + } + if request == nil { + return nil, providerUtils.NewBifrostOperationError("file delete request is nil", nil) + } + if strings.TrimSpace(request.FileID) == "" { + return nil, providerUtils.NewBifrostOperationError("file_id is required", nil) + } + if len(keys) == 0 { + keys = []schemas.Key{{}} + } + + var lastErr *schemas.BifrostError + for _, key := range keys { + response, bifrostErr := provider.fileDeleteWithRefresh(ctx, key, request, false) + if isGigaChatUnauthorizedError(bifrostErr) { + response, bifrostErr = provider.fileDeleteWithRefresh(ctx, key, request, true) + } + if bifrostErr == nil { + return response, nil + } + lastErr = bifrostErr + } + return nil, lastErr +} + +// FileContent downloads GigaChat file content. +func (provider *GigaChatProvider) FileContent(ctx *schemas.BifrostContext, keys []schemas.Key, request *schemas.BifrostFileContentRequest) (*schemas.BifrostFileContentResponse, *schemas.BifrostError) { + if err := providerUtils.CheckOperationAllowed(schemas.GigaChat, provider.customProviderConfig, schemas.FileContentRequest); err != nil { + return nil, err + } + if request == nil { + return nil, providerUtils.NewBifrostOperationError("file content request is nil", nil) + } + if strings.TrimSpace(request.FileID) == "" { + return nil, providerUtils.NewBifrostOperationError("file_id is required", nil) + } + if len(keys) == 0 { + keys = []schemas.Key{{}} + } + + var lastErr *schemas.BifrostError + for _, key := range keys { + response, bifrostErr := provider.fileContentWithRefresh(ctx, key, request, false) + if isGigaChatUnauthorizedError(bifrostErr) { + response, bifrostErr = provider.fileContentWithRefresh(ctx, key, request, true) + } + if bifrostErr == nil { + return response, nil + } + lastErr = bifrostErr + } + return nil, lastErr +} + +// CachedContentCreate is not supported by the GigaChat provider skeleton. +func (provider *GigaChatProvider) CachedContentCreate(_ *schemas.BifrostContext, _ schemas.Key, _ *schemas.BifrostCachedContentCreateRequest) (*schemas.BifrostCachedContentCreateResponse, *schemas.BifrostError) { + return nil, provider.unsupported(schemas.CachedContentCreateRequest) +} + +// CachedContentList is not supported by the GigaChat provider skeleton. +func (provider *GigaChatProvider) CachedContentList(_ *schemas.BifrostContext, _ []schemas.Key, _ *schemas.BifrostCachedContentListRequest) (*schemas.BifrostCachedContentListResponse, *schemas.BifrostError) { + return nil, provider.unsupported(schemas.CachedContentListRequest) +} + +// CachedContentRetrieve is not supported by the GigaChat provider skeleton. +func (provider *GigaChatProvider) CachedContentRetrieve(_ *schemas.BifrostContext, _ []schemas.Key, _ *schemas.BifrostCachedContentRetrieveRequest) (*schemas.BifrostCachedContentRetrieveResponse, *schemas.BifrostError) { + return nil, provider.unsupported(schemas.CachedContentRetrieveRequest) +} + +// CachedContentUpdate is not supported by the GigaChat provider skeleton. +func (provider *GigaChatProvider) CachedContentUpdate(_ *schemas.BifrostContext, _ []schemas.Key, _ *schemas.BifrostCachedContentUpdateRequest) (*schemas.BifrostCachedContentUpdateResponse, *schemas.BifrostError) { + return nil, provider.unsupported(schemas.CachedContentUpdateRequest) +} + +// CachedContentDelete is not supported by the GigaChat provider skeleton. +func (provider *GigaChatProvider) CachedContentDelete(_ *schemas.BifrostContext, _ []schemas.Key, _ *schemas.BifrostCachedContentDeleteRequest) (*schemas.BifrostCachedContentDeleteResponse, *schemas.BifrostError) { + return nil, provider.unsupported(schemas.CachedContentDeleteRequest) +} + +// ContainerCreate is not supported by the GigaChat provider skeleton. +func (provider *GigaChatProvider) ContainerCreate(_ *schemas.BifrostContext, _ schemas.Key, _ *schemas.BifrostContainerCreateRequest) (*schemas.BifrostContainerCreateResponse, *schemas.BifrostError) { + return nil, provider.unsupported(schemas.ContainerCreateRequest) +} + +// ContainerList is not supported by the GigaChat provider skeleton. +func (provider *GigaChatProvider) ContainerList(_ *schemas.BifrostContext, _ []schemas.Key, _ *schemas.BifrostContainerListRequest) (*schemas.BifrostContainerListResponse, *schemas.BifrostError) { + return nil, provider.unsupported(schemas.ContainerListRequest) +} + +// ContainerRetrieve is not supported by the GigaChat provider skeleton. +func (provider *GigaChatProvider) ContainerRetrieve(_ *schemas.BifrostContext, _ []schemas.Key, _ *schemas.BifrostContainerRetrieveRequest) (*schemas.BifrostContainerRetrieveResponse, *schemas.BifrostError) { + return nil, provider.unsupported(schemas.ContainerRetrieveRequest) +} + +// ContainerDelete is not supported by the GigaChat provider skeleton. +func (provider *GigaChatProvider) ContainerDelete(_ *schemas.BifrostContext, _ []schemas.Key, _ *schemas.BifrostContainerDeleteRequest) (*schemas.BifrostContainerDeleteResponse, *schemas.BifrostError) { + return nil, provider.unsupported(schemas.ContainerDeleteRequest) +} + +// ContainerFileCreate is not supported by the GigaChat provider skeleton. +func (provider *GigaChatProvider) ContainerFileCreate(_ *schemas.BifrostContext, _ schemas.Key, _ *schemas.BifrostContainerFileCreateRequest) (*schemas.BifrostContainerFileCreateResponse, *schemas.BifrostError) { + return nil, provider.unsupported(schemas.ContainerFileCreateRequest) +} + +// ContainerFileList is not supported by the GigaChat provider skeleton. +func (provider *GigaChatProvider) ContainerFileList(_ *schemas.BifrostContext, _ []schemas.Key, _ *schemas.BifrostContainerFileListRequest) (*schemas.BifrostContainerFileListResponse, *schemas.BifrostError) { + return nil, provider.unsupported(schemas.ContainerFileListRequest) +} + +// ContainerFileRetrieve is not supported by the GigaChat provider skeleton. +func (provider *GigaChatProvider) ContainerFileRetrieve(_ *schemas.BifrostContext, _ []schemas.Key, _ *schemas.BifrostContainerFileRetrieveRequest) (*schemas.BifrostContainerFileRetrieveResponse, *schemas.BifrostError) { + return nil, provider.unsupported(schemas.ContainerFileRetrieveRequest) +} + +// ContainerFileContent is not supported by the GigaChat provider skeleton. +func (provider *GigaChatProvider) ContainerFileContent(_ *schemas.BifrostContext, _ []schemas.Key, _ *schemas.BifrostContainerFileContentRequest) (*schemas.BifrostContainerFileContentResponse, *schemas.BifrostError) { + return nil, provider.unsupported(schemas.ContainerFileContentRequest) +} + +// ContainerFileDelete is not supported by the GigaChat provider skeleton. +func (provider *GigaChatProvider) ContainerFileDelete(_ *schemas.BifrostContext, _ []schemas.Key, _ *schemas.BifrostContainerFileDeleteRequest) (*schemas.BifrostContainerFileDeleteResponse, *schemas.BifrostError) { + return nil, provider.unsupported(schemas.ContainerFileDeleteRequest) +} + +// Passthrough is not supported by the GigaChat provider skeleton. +func (provider *GigaChatProvider) Passthrough(_ *schemas.BifrostContext, _ schemas.Key, _ *schemas.BifrostPassthroughRequest) (*schemas.BifrostPassthroughResponse, *schemas.BifrostError) { + return nil, provider.unsupported(schemas.PassthroughRequest) +} + +// PassthroughStream is not supported by the GigaChat provider skeleton. +func (provider *GigaChatProvider) PassthroughStream(_ *schemas.BifrostContext, _ schemas.PostHookRunner, _ func(context.Context), _ schemas.Key, _ *schemas.BifrostPassthroughRequest) (chan *schemas.BifrostStreamChunk, *schemas.BifrostError) { + return nil, provider.unsupported(schemas.PassthroughStreamRequest) +} diff --git a/core/providers/gigachat/models.go b/core/providers/gigachat/models.go new file mode 100644 index 00000000000..49b7fd30350 --- /dev/null +++ b/core/providers/gigachat/models.go @@ -0,0 +1,79 @@ +package gigachat + +import ( + "strings" + + providerUtils "github.com/maximhq/bifrost/core/providers/utils" + schemas "github.com/maximhq/bifrost/core/schemas" +) + +// ToBifrostListModelsResponse converts GigaChat model metadata to Bifrost format. +func (response *GigaChatListModelsResponse) ToBifrostListModelsResponse(providerKey schemas.ModelProvider, allowedModels schemas.WhiteList, blacklistedModels schemas.BlackList, aliases schemas.KeyAliases, unfiltered bool) *schemas.BifrostListModelsResponse { + if response == nil { + return nil + } + + bifrostResponse := &schemas.BifrostListModelsResponse{ + Data: make([]schemas.Model, 0, len(response.Data)), + } + + pipeline := &providerUtils.ListModelsPipeline{ + AllowedModels: allowedModels, + BlacklistedModels: blacklistedModels, + Aliases: aliases, + Unfiltered: unfiltered, + ProviderKey: providerKey, + MatchFns: providerUtils.DefaultMatchFns(), + } + if pipeline.ShouldEarlyExit() { + return bifrostResponse + } + + included := make(map[string]bool) + for _, model := range response.Data { + modelID := strings.TrimSpace(model.ID) + if modelID == "" { + continue + } + + for _, result := range pipeline.FilterModel(modelID) { + entry := schemas.Model{ + ID: string(providerKey) + "/" + result.ResolvedID, + OwnedBy: stringPtrIfNotEmpty(model.OwnedBy), + SupportedMethods: toGigaChatSupportedMethods(model.Type), + } + if result.AliasValue != "" { + entry.Alias = schemas.Ptr(result.AliasValue) + } + bifrostResponse.Data = append(bifrostResponse.Data, entry) + included[strings.ToLower(result.ResolvedID)] = true + } + } + + bifrostResponse.Data = append(bifrostResponse.Data, pipeline.BackfillModels(included)...) + return bifrostResponse +} + +func stringPtrIfNotEmpty(value string) *string { + value = strings.TrimSpace(value) + if value == "" { + return nil + } + return schemas.Ptr(value) +} + +func toGigaChatSupportedMethods(modelType string) []string { + switch strings.ToLower(strings.TrimSpace(modelType)) { + case "chat": + return []string{ + string(schemas.ChatCompletionRequest), + string(schemas.ChatCompletionStreamRequest), + string(schemas.ResponsesRequest), + string(schemas.ResponsesStreamRequest), + } + case "embedder", "embedding", "embeddings": + return []string{string(schemas.EmbeddingRequest)} + default: + return nil + } +} diff --git a/core/providers/gigachat/responses.go b/core/providers/gigachat/responses.go new file mode 100644 index 00000000000..30547aa8294 --- /dev/null +++ b/core/providers/gigachat/responses.go @@ -0,0 +1,1799 @@ +package gigachat + +import ( + "encoding/base64" + "encoding/json" + "fmt" + "sort" + "strconv" + "strings" + + "github.com/bytedance/sonic" + schemas "github.com/maximhq/bifrost/core/schemas" +) + +const ( + gigaChatResponsesRoleReasoning = "reasoning" + gigaChatResponsesGeneratedCallIDPrefix = "gigachat_call_" + gigaChatResponsesGeneratedCallIDVersion = "v1" +) + +// ToGigaChatResponsesRequest converts a Bifrost Responses request to GigaChat v2 chat completions format. +func ToGigaChatResponsesRequest(bifrostReq *schemas.BifrostResponsesRequest) (*GigaChatResponsesRequest, error) { + if bifrostReq == nil { + return nil, fmt.Errorf("bifrost responses request is nil") + } + if strings.TrimSpace(bifrostReq.Model) == "" { + return nil, fmt.Errorf("model is required") + } + + messages := make([]GigaChatResponsesMessage, 0, len(bifrostReq.Input)+1) + if bifrostReq.Params != nil && bifrostReq.Params.Instructions != nil && strings.TrimSpace(*bifrostReq.Params.Instructions) != "" { + messages = append(messages, GigaChatResponsesMessage{ + Role: string(schemas.ResponsesInputMessageRoleSystem), + Content: []GigaChatResponsesContentPart{{ + Text: schemas.Ptr(*bifrostReq.Params.Instructions), + }}, + }) + } + + functionCallNamesByID := collectGigaChatResponsesFunctionCallNames(bifrostReq.Input) + for index, message := range bifrostReq.Input { + convertedMessages, err := toGigaChatResponsesMessages(message, functionCallNamesByID) + if err != nil { + return nil, fmt.Errorf("input[%d]: %w", index, err) + } + messages = append(messages, convertedMessages...) + } + if len(messages) == 0 { + return nil, fmt.Errorf("messages are required") + } + + gigaChatReq := &GigaChatResponsesRequest{ + Model: bifrostReq.Model, + Messages: messages, + } + params := bifrostReq.Params + if params == nil { + params = &schemas.ResponsesParameters{} + } + + if unsupportedParams := unsupportedGigaChatResponsesParams(params); len(unsupportedParams) > 0 { + return nil, fmt.Errorf("GigaChat Responses do not support parameter(s): %s", strings.Join(unsupportedParams, ", ")) + } + if err := applyGigaChatResponsesParams(gigaChatReq, params); err != nil { + return nil, err + } + if hasGigaChatResponsesThreadID(params) { + gigaChatReq.Model = "" + } + return gigaChatReq, nil +} + +// ToGigaChatResponsesStreamRequest converts a Bifrost Responses request to a streaming GigaChat v2 request. +func ToGigaChatResponsesStreamRequest(bifrostReq *schemas.BifrostResponsesRequest) (*GigaChatResponsesRequest, error) { + gigaChatReq, err := ToGigaChatResponsesRequest(bifrostReq) + if err != nil { + return nil, err + } + gigaChatReq.Stream = schemas.Ptr(true) + return gigaChatReq, nil +} + +// ToBifrostResponsesResponse converts a GigaChat v2 chat completions response to Bifrost Responses format. +func ToBifrostResponsesResponse(providerName schemas.ModelProvider, response *GigaChatResponsesResponse) *schemas.BifrostResponsesResponse { + if response == nil { + return nil + } + + outputCapacity := len(response.Messages) + if outputCapacity == 0 { + outputCapacity = len(response.Choices) + } + output := make([]schemas.ResponsesMessage, 0, outputCapacity) + var status *string + var incompleteDetails *schemas.ResponsesResponseIncompleteDetails + var stopReason *string + functionCallIDs := newGigaChatResponsesCallIDTracker() + applyFinishReason := func(finishReason *string) bool { + if mappedStatus, mappedIncompleteDetails, mappedStopReason := toBifrostGigaChatResponsesStatus(finishReason); mappedStatus != nil { + status = mappedStatus + incompleteDetails = mappedIncompleteDetails + stopReason = mappedStopReason + return *mappedStatus == "incomplete" + } + return false + } + + if len(response.Messages) > 0 { + for _, message := range response.Messages { + output = append(output, toBifrostGigaChatResponsesMessageOutput(message, response.MessageID, response.ToolsStateID, functionCallIDs)...) + + finishReason := message.FinishReason + if finishReason == nil { + finishReason = response.FinishReason + } + if applyFinishReason(finishReason) { + break + } + } + } else { + for _, choice := range response.Choices { + output = append(output, toBifrostGigaChatResponsesChoiceOutput(choice, response.MessageID, response.ToolsStateID, functionCallIDs)...) + + finishReason := choice.FinishReason + if finishReason == nil && choice.Message != nil { + finishReason = choice.Message.FinishReason + } + if applyFinishReason(finishReason) { + break + } + } + } + if status == nil { + applyFinishReason(response.FinishReason) + } + + createdAt := response.CreatedAt + if createdAt == 0 { + createdAt = response.Created + } + responseID := toBifrostGigaChatResponsesResponseID(response) + + bifrostResponse := &schemas.BifrostResponsesResponse{ + Object: "response", + CreatedAt: createdAt, + Conversation: toBifrostGigaChatResponsesConversation(response.ThreadID), + Model: response.Model, + Output: output, + Status: status, + IncompleteDetails: incompleteDetails, + StopReason: stopReason, + ExtraFields: schemas.BifrostResponseExtraFields{ + Provider: providerName, + }, + } + if strings.TrimSpace(responseID) != "" { + bifrostResponse.ID = &responseID + } + if usage := toBifrostGigaChatUsage(response.Usage); usage != nil { + bifrostResponse.Usage = usage.ToResponsesResponseUsage() + } + bifrostResponse.ProviderExtraFields = toBifrostGigaChatResponsesProviderExtraFields(response) + + return bifrostResponse +} + +// ToBifrostResponsesStreamResponse converts a GigaChat v2 SSE chunk to Bifrost Responses stream events. +func ToBifrostResponsesStreamResponse(providerName schemas.ModelProvider, response *GigaChatResponsesResponse, state *schemas.ChatToResponsesStreamState) []*schemas.BifrostResponsesStreamResponse { + if response == nil || state == nil { + return nil + } + + if response.CreatedAt != 0 && !state.HasEmittedCreated { + state.CreatedAt = response.CreatedAt + } else if response.Created != 0 && !state.HasEmittedCreated { + state.CreatedAt = response.Created + } + + chatResponse := toBifrostGigaChatResponsesChatStreamResponse(providerName, response) + if chatResponse == nil { + return nil + } + ensureGigaChatResponsesStreamLifecycleRole(chatResponse, state) + + events := chatResponse.ToBifrostResponsesStreamResponse(state) + conversation := toBifrostGigaChatResponsesConversation(response.ThreadID) + for _, event := range events { + if event == nil { + continue + } + applyGigaChatResponsesStreamExtraFields(event, providerName) + if event.Response != nil && event.Response.Conversation == nil { + event.Response.Conversation = conversation + } + } + return events +} + +func applyGigaChatResponsesStreamExtraFields(event *schemas.BifrostResponsesStreamResponse, providerName schemas.ModelProvider) { + if event == nil { + return + } + event.ExtraFields.Provider = providerName + event.ExtraFields.RequestType = schemas.ResponsesStreamRequest + if event.Response == nil { + return + } + event.Response.ExtraFields.Provider = providerName + event.Response.ExtraFields.RequestType = schemas.ResponsesStreamRequest +} + +func toBifrostGigaChatResponsesResponseID(response *GigaChatResponsesResponse) string { + if response == nil { + return "" + } + if responseID := strings.TrimSpace(response.ID); responseID != "" { + return responseID + } + if response.ThreadID != nil { + if threadID := strings.TrimSpace(*response.ThreadID); threadID != "" { + return threadID + } + } + if response.MessageID != nil { + return strings.TrimSpace(*response.MessageID) + } + return "" +} + +func toBifrostGigaChatResponsesConversation(threadID *string) *schemas.ResponsesResponseConversation { + if threadID == nil || strings.TrimSpace(*threadID) == "" { + return nil + } + return &schemas.ResponsesResponseConversation{ + ResponsesResponseConversationStruct: &schemas.ResponsesResponseConversationStruct{ + ID: strings.TrimSpace(*threadID), + }, + } +} + +func toBifrostGigaChatResponsesChatStreamResponse(providerName schemas.ModelProvider, response *GigaChatResponsesResponse) *schemas.BifrostChatResponse { + if response == nil { + return nil + } + + createdAt := response.CreatedAt + if createdAt == 0 { + createdAt = response.Created + } + responseID := toBifrostGigaChatResponsesResponseID(response) + + streamResponse := &GigaChatChatStreamResponse{ + ID: responseID, + Created: createdAt, + Model: response.Model, + Object: response.Object, + SystemFingerprint: response.SystemFingerprint, + Usage: response.Usage, + ExtraParams: response.ExtraParams, + } + + if len(response.Messages) > 0 { + streamResponse.Choices = make([]GigaChatChatStreamChoice, 0, len(response.Messages)) + for index, message := range response.Messages { + finishReason := message.FinishReason + if finishReason == nil { + finishReason = response.FinishReason + } + streamResponse.Choices = append(streamResponse.Choices, GigaChatChatStreamChoice{ + Index: index, + Delta: toGigaChatResponsesMessageStreamDelta(&message, response.ToolsStateID), + FinishReason: finishReason, + }) + } + return ToBifrostChatStreamResponse(providerName, streamResponse) + } + + if len(response.Choices) > 0 { + streamResponse.Choices = make([]GigaChatChatStreamChoice, 0, len(response.Choices)) + for _, choice := range response.Choices { + delta := choice.Delta + if delta == nil && choice.Message != nil { + delta = toGigaChatResponsesMessageStreamDelta(choice.Message, response.ToolsStateID) + } + streamResponse.Choices = append(streamResponse.Choices, GigaChatChatStreamChoice{ + Index: choice.Index, + Delta: delta, + FinishReason: choice.FinishReason, + LogProbs: choice.LogProbs, + }) + } + return ToBifrostChatStreamResponse(providerName, streamResponse) + } + + if response.FinishReason != nil { + streamResponse.Choices = []GigaChatChatStreamChoice{{ + Index: 0, + Delta: &GigaChatChatStreamDelta{}, + FinishReason: response.FinishReason, + }} + return ToBifrostChatStreamResponse(providerName, streamResponse) + } + + return nil +} + +func toGigaChatResponsesMessageStreamDelta(message *GigaChatResponsesMessage, fallbackToolsStateID *string) *GigaChatChatStreamDelta { + if message == nil { + return &GigaChatChatStreamDelta{} + } + + delta := &GigaChatChatStreamDelta{} + if strings.TrimSpace(message.Role) != "" { + role := message.Role + if isGigaChatResponsesReasoningRole(role) { + role = string(schemas.ChatMessageRoleAssistant) + } + delta.Role = &role + } + + var textBuilder strings.Builder + var functionCall *GigaChatResponsesFunctionCall + for index := range message.Content { + part := message.Content[index] + if part.Text != nil { + textBuilder.WriteString(*part.Text) + } + if functionCall == nil && part.FunctionCall != nil { + functionCall = part.FunctionCall + } + } + if text := textBuilder.String(); text != "" { + if isGigaChatResponsesReasoningRole(message.Role) { + delta.Reasoning = &text + } else { + delta.Content = &text + } + } + if functionCall == nil { + functionCall = message.FunctionCall + } + if functionCall != nil { + delta.FunctionCall = toGigaChatLegacyFunctionCall(functionCall) + delta.FunctionsStateID = toGigaChatResponsesMessageToolStateID(*message) + if delta.FunctionsStateID == nil { + delta.FunctionsStateID = fallbackToolsStateID + } + } + + return delta +} + +func toGigaChatLegacyFunctionCall(functionCall *GigaChatResponsesFunctionCall) *GigaChatFunctionCall { + if functionCall == nil { + return nil + } + arguments := json.RawMessage(stringifyGigaChatResponsesPayload(functionCall.Arguments)) + return &GigaChatFunctionCall{ + Name: toBifrostGigaChatResponsesFunctionName(functionCall.Name), + Arguments: arguments, + } +} + +func ensureGigaChatResponsesStreamLifecycleRole(response *schemas.BifrostChatResponse, state *schemas.ChatToResponsesStreamState) { + if response == nil || state == nil || state.HasEmittedCreated || len(response.Choices) == 0 { + return + } + choice := response.Choices[0] + if choice.ChatStreamResponseChoice == nil || choice.ChatStreamResponseChoice.Delta == nil { + return + } + delta := choice.ChatStreamResponseChoice.Delta + if delta.Role != nil { + return + } + hasContent := delta.Content != nil && *delta.Content != "" + if hasContent || len(delta.ToolCalls) > 0 { + role := string(schemas.ChatMessageRoleAssistant) + delta.Role = &role + } +} + +func updateGigaChatResponsesStreamUsage(target *schemas.BifrostLLMUsage, source *schemas.BifrostLLMUsage) { + if target == nil || source == nil { + return + } + if source.PromptTokens > target.PromptTokens { + target.PromptTokens = source.PromptTokens + } + if source.CompletionTokens > target.CompletionTokens { + target.CompletionTokens = source.CompletionTokens + } + if source.TotalTokens > target.TotalTokens { + target.TotalTokens = source.TotalTokens + } + if calculatedTotal := target.PromptTokens + target.CompletionTokens; calculatedTotal > target.TotalTokens { + target.TotalTokens = calculatedTotal + } + if source.PromptTokensDetails != nil { + target.PromptTokensDetails = source.PromptTokensDetails + } + if source.CompletionTokensDetails != nil { + target.CompletionTokensDetails = source.CompletionTokensDetails + } + if source.Cost != nil { + target.Cost = source.Cost + } +} + +func toBifrostGigaChatResponsesChoiceOutput(choice GigaChatResponsesChoice, fallbackMessageID *string, fallbackToolsStateID *string, functionCallIDs *gigaChatResponsesCallIDTracker) []schemas.ResponsesMessage { + if choice.Message == nil { + return nil + } + + return toBifrostGigaChatResponsesMessageOutput(*choice.Message, fallbackMessageID, fallbackToolsStateID, functionCallIDs) +} + +func toBifrostGigaChatResponsesMessageOutput(message GigaChatResponsesMessage, fallbackMessageID *string, fallbackToolsStateID *string, functionCallIDs *gigaChatResponsesCallIDTracker) []schemas.ResponsesMessage { + messageID := message.MessageID + if messageID == nil || strings.TrimSpace(*messageID) == "" { + messageID = fallbackMessageID + } + if isGigaChatResponsesReasoningRole(message.Role) { + return toBifrostGigaChatResponsesReasoningOutput(message, messageID) + } + + toolsStateID := toGigaChatResponsesMessageToolStateID(message) + if toolsStateID == nil || strings.TrimSpace(*toolsStateID) == "" { + toolsStateID = fallbackToolsStateID + } + output := make([]schemas.ResponsesMessage, 0, len(message.Content)+1) + contentBlocks := make([]schemas.ResponsesMessageContentBlock, 0, len(message.Content)) + hasFunctionCall := false + for index, part := range message.Content { + sourceRefs := toBifrostGigaChatResponsesInlineSources(part.InlineData) + if part.Text != nil { + annotations := toBifrostGigaChatResponsesSourceAnnotations(*part.Text, sourceRefs) + if annotations == nil { + annotations = []schemas.ResponsesOutputMessageContentTextAnnotation{} + } + contentBlocks = append(contentBlocks, schemas.ResponsesMessageContentBlock{ + Type: schemas.ResponsesOutputMessageContentTypeText, + Text: part.Text, + ResponsesOutputMessageContentText: &schemas.ResponsesOutputMessageContentText{ + Annotations: annotations, + LogProbs: []schemas.ResponsesOutputMessageContentTextLogProb{}, + }, + }) + } + if webSearchCall := toBifrostGigaChatResponsesWebSearchCall(messageID, index, sourceRefs); webSearchCall != nil { + output = append(output, *webSearchCall) + } + for fileIndex, file := range part.Files { + if imageCall := toBifrostGigaChatResponsesImageGenerationCall(messageID, index, fileIndex, file); imageCall != nil { + output = append(output, *imageCall) + } + } + if part.FunctionCall != nil { + hasFunctionCall = true + if toolCall := toBifrostGigaChatResponsesFunctionCall(messageID, toolsStateID, index, functionCallIDs, part.FunctionCall); toolCall != nil { + output = append(output, *toolCall) + } + } + if part.FunctionResult != nil { + if toolResult := toBifrostGigaChatResponsesFunctionResult(messageID, toolsStateID, index, functionCallIDs, part.FunctionResult); toolResult != nil { + output = append(output, *toolResult) + } + } + } + if len(contentBlocks) > 0 { + messageType := schemas.ResponsesMessageTypeMessage + role := schemas.ResponsesInputMessageRoleAssistant + if strings.TrimSpace(message.Role) != "" { + role = schemas.ResponsesMessageRoleType(message.Role) + } + output = append([]schemas.ResponsesMessage{{ + ID: messageID, + Type: &messageType, + Role: &role, + Status: schemas.Ptr("completed"), + Content: &schemas.ResponsesMessageContent{ + ContentBlocks: contentBlocks, + }, + }}, output...) + } + if !hasFunctionCall && message.FunctionCall != nil { + if toolCall := toBifrostGigaChatResponsesFunctionCall(messageID, toolsStateID, 0, functionCallIDs, message.FunctionCall); toolCall != nil { + output = append(output, *toolCall) + } + } + return output +} + +type gigaChatResponsesInlineSource struct { + Key string + Order int + HasOrder bool + URL string + Title string +} + +func toBifrostGigaChatResponsesWebSearchCall(messageID *string, partIndex int, sources []gigaChatResponsesInlineSource) *schemas.ResponsesMessage { + if len(sources) == 0 { + return nil + } + + actionSources := make([]schemas.ResponsesWebSearchToolCallActionSearchSource, 0, len(sources)) + for _, source := range sources { + if strings.TrimSpace(source.URL) == "" { + continue + } + actionSource := schemas.ResponsesWebSearchToolCallActionSearchSource{ + Type: "url", + URL: source.URL, + } + if strings.TrimSpace(source.Title) != "" { + actionSource.Title = schemas.Ptr(source.Title) + } + actionSources = append(actionSources, actionSource) + } + if len(actionSources) == 0 { + return nil + } + + itemID := toBifrostGigaChatResponsesFileItemID("ws", messageID, partIndex, 0) + messageType := schemas.ResponsesMessageTypeWebSearchCall + + return &schemas.ResponsesMessage{ + ID: &itemID, + Type: &messageType, + Status: schemas.Ptr("completed"), + ResponsesToolMessage: &schemas.ResponsesToolMessage{ + Action: &schemas.ResponsesToolMessageActionStruct{ + ResponsesWebSearchToolCallAction: &schemas.ResponsesWebSearchToolCallAction{ + Type: "search", + Sources: actionSources, + }, + }, + }, + } +} + +func toBifrostGigaChatResponsesSourceAnnotations(text string, sources []gigaChatResponsesInlineSource) []schemas.ResponsesOutputMessageContentTextAnnotation { + if len(sources) == 0 { + return nil + } + + citedSources, startIndex, endIndex := toBifrostGigaChatResponsesCitedSources(text) + useAllSources := len(citedSources) == 0 + if useAllSources && text != "" { + startIndex = schemas.Ptr(0) + endIndex = schemas.Ptr(len(text)) + } + + annotations := make([]schemas.ResponsesOutputMessageContentTextAnnotation, 0, len(sources)) + for _, source := range sources { + if strings.TrimSpace(source.URL) == "" { + continue + } + if !useAllSources { + if _, ok := citedSources[source.Key]; !ok { + continue + } + } + + annotation := schemas.ResponsesOutputMessageContentTextAnnotation{ + Type: "url_citation", + URL: schemas.Ptr(source.URL), + StartIndex: startIndex, + EndIndex: endIndex, + } + if strings.TrimSpace(source.Title) != "" { + annotation.Title = schemas.Ptr(source.Title) + } + annotations = append(annotations, annotation) + } + if len(annotations) == 0 { + return nil + } + return annotations +} + +func toBifrostGigaChatResponsesCitedSources(text string) (map[string]struct{}, *int, *int) { + const markerPrefix = "[sources=[" + + start := strings.LastIndex(text, markerPrefix) + if start < 0 { + return nil, nil, nil + } + + sourceListStart := start + len(markerPrefix) + remainder := text[sourceListStart:] + sourceListEnd := strings.Index(remainder, "]]") + markerSuffixLen := 2 + if sourceListEnd < 0 { + sourceListEnd = strings.Index(remainder, "]") + markerSuffixLen = 1 + } + if sourceListEnd < 0 { + return nil, nil, nil + } + + citedSources := make(map[string]struct{}) + for _, rawSource := range strings.Split(remainder[:sourceListEnd], ",") { + sourceKey := strings.TrimSpace(rawSource) + if sourceKey != "" { + citedSources[sourceKey] = struct{}{} + } + } + if len(citedSources) == 0 { + return nil, nil, nil + } + + end := sourceListStart + sourceListEnd + markerSuffixLen + return citedSources, &start, &end +} + +func toBifrostGigaChatResponsesInlineSources(inlineData map[string]interface{}) []gigaChatResponsesInlineSource { + if len(inlineData) == 0 { + return nil + } + + rawSources, ok := inlineData["sources"] + if !ok { + return nil + } + + sourcesMap, ok := schemas.SafeExtractOrderedMap(rawSources) + if !ok || sourcesMap.Len() == 0 { + return nil + } + + sources := make([]gigaChatResponsesInlineSource, 0, sourcesMap.Len()) + sourcesMap.Range(func(key string, value interface{}) bool { + if source, ok := toBifrostGigaChatResponsesInlineSource(key, value); ok { + sources = append(sources, source) + } + return true + }) + if len(sources) == 0 { + return nil + } + + sort.SliceStable(sources, func(i, j int) bool { + left := sources[i] + right := sources[j] + if left.HasOrder && right.HasOrder { + return left.Order < right.Order + } + if left.HasOrder != right.HasOrder { + return left.HasOrder + } + return left.Key < right.Key + }) + return sources +} + +func toBifrostGigaChatResponsesInlineSource(key string, value interface{}) (gigaChatResponsesInlineSource, bool) { + sourceMap, ok := schemas.SafeExtractOrderedMap(value) + if !ok { + return gigaChatResponsesInlineSource{}, false + } + + rawURL, ok := sourceMap.Get("url") + if !ok { + return gigaChatResponsesInlineSource{}, false + } + url, ok := schemas.SafeExtractString(rawURL) + if !ok || strings.TrimSpace(url) == "" { + return gigaChatResponsesInlineSource{}, false + } + + source := gigaChatResponsesInlineSource{ + Key: strings.TrimSpace(key), + URL: strings.TrimSpace(url), + } + if order, err := strconv.Atoi(source.Key); err == nil { + source.Order = order + source.HasOrder = true + } + if rawTitle, ok := sourceMap.Get("title"); ok { + if title, ok := schemas.SafeExtractString(rawTitle); ok { + source.Title = strings.TrimSpace(title) + } + } + return source, true +} + +func toBifrostGigaChatResponsesImageGenerationCall(messageID *string, partIndex int, fileIndex int, file GigaChatResponsesContentFile) *schemas.ResponsesMessage { + fileID := strings.TrimSpace(file.ID) + if fileID == "" || !isGigaChatResponsesImageFile(file) { + return nil + } + + itemID := toBifrostGigaChatResponsesFileItemID("ig", messageID, partIndex, fileIndex) + messageType := schemas.ResponsesMessageTypeImageGenerationCall + + return &schemas.ResponsesMessage{ + ID: &itemID, + Type: &messageType, + Status: schemas.Ptr("completed"), + ResponsesToolMessage: &schemas.ResponsesToolMessage{ + ResponsesImageGenerationCall: &schemas.ResponsesImageGenerationCall{ + Result: fileID, + }, + }, + } +} + +func isGigaChatResponsesImageFile(file GigaChatResponsesContentFile) bool { + if file.Target != nil && strings.EqualFold(strings.TrimSpace(*file.Target), "image") { + return true + } + return file.MIME != nil && strings.HasPrefix(strings.ToLower(strings.TrimSpace(*file.MIME)), "image/") +} + +func toBifrostGigaChatResponsesReasoningOutput(message GigaChatResponsesMessage, messageID *string) []schemas.ResponsesMessage { + var textBuilder strings.Builder + for _, part := range message.Content { + if part.Text != nil { + textBuilder.WriteString(*part.Text) + } + } + + reasoningText := textBuilder.String() + if strings.TrimSpace(reasoningText) == "" { + return nil + } + + messageType := schemas.ResponsesMessageTypeReasoning + role := schemas.ResponsesInputMessageRoleAssistant + itemID := toBifrostGigaChatResponsesReasoningItemID(messageID) + return []schemas.ResponsesMessage{{ + ID: itemID, + Type: &messageType, + Role: &role, + Status: schemas.Ptr("completed"), + ResponsesReasoning: &schemas.ResponsesReasoning{ + Summary: []schemas.ResponsesReasoningSummary{{ + Type: schemas.ResponsesReasoningContentBlockTypeSummaryText, + Text: reasoningText, + }}, + }, + }} +} + +func toBifrostGigaChatResponsesFunctionCall(messageID *string, toolsStateID *string, index int, functionCallIDs *gigaChatResponsesCallIDTracker, functionCall *GigaChatResponsesFunctionCall) *schemas.ResponsesMessage { + if functionCall == nil || strings.TrimSpace(functionCall.Name) == "" { + return nil + } + + itemID := toBifrostGigaChatResponsesItemID("fc", messageID, index) + callID := functionCallIDs.FunctionCallID(toolsStateID, itemID, functionCall.Name) + arguments := stringifyGigaChatResponsesPayload(functionCall.Arguments) + messageType := schemas.ResponsesMessageTypeFunctionCall + role := schemas.ResponsesInputMessageRoleAssistant + return &schemas.ResponsesMessage{ + ID: &itemID, + Type: &messageType, + Role: &role, + Status: schemas.Ptr("completed"), + ResponsesToolMessage: &schemas.ResponsesToolMessage{ + CallID: &callID, + Name: schemas.Ptr(toBifrostGigaChatResponsesFunctionName(functionCall.Name)), + Arguments: &arguments, + }, + } +} + +func toBifrostGigaChatResponsesFunctionResult(messageID *string, toolsStateID *string, index int, functionCallIDs *gigaChatResponsesCallIDTracker, functionResult *GigaChatResponsesFunctionResult) *schemas.ResponsesMessage { + if functionResult == nil || strings.TrimSpace(functionResult.Name) == "" { + return nil + } + + itemID := toBifrostGigaChatResponsesItemID("fr", messageID, index) + callID := functionCallIDs.FunctionResultID(toolsStateID, itemID, functionResult.Name) + output := stringifyGigaChatResponsesPayload(functionResult.Result) + messageType := schemas.ResponsesMessageTypeFunctionCallOutput + return &schemas.ResponsesMessage{ + ID: &itemID, + Type: &messageType, + Status: schemas.Ptr("completed"), + ResponsesToolMessage: &schemas.ResponsesToolMessage{ + CallID: &callID, + Name: schemas.Ptr(toBifrostGigaChatResponsesFunctionName(functionResult.Name)), + Output: &schemas.ResponsesToolMessageOutputStruct{ + ResponsesToolCallOutputStr: &output, + }, + }, + } +} + +type gigaChatResponsesCallIDTracker struct { + counts map[string]int + pendingByInvocation map[string][]string +} + +func newGigaChatResponsesCallIDTracker() *gigaChatResponsesCallIDTracker { + return &gigaChatResponsesCallIDTracker{ + counts: make(map[string]int), + pendingByInvocation: make(map[string][]string), + } +} + +func (tracker *gigaChatResponsesCallIDTracker) FunctionCallID(toolsStateID *string, fallback string, name string) string { + trimmed := trimStringPtr(toolsStateID) + if trimmed == "" { + return fallback + } + publicID := tracker.nextCallID(trimmed, name) + if tracker != nil { + key := tracker.invocationKey(trimmed, name) + tracker.pendingByInvocation[key] = append(tracker.pendingByInvocation[key], publicID) + } + return publicID +} + +func (tracker *gigaChatResponsesCallIDTracker) FunctionResultID(toolsStateID *string, fallback string, name string) string { + trimmed := trimStringPtr(toolsStateID) + if trimmed == "" { + return fallback + } + if tracker != nil { + key := tracker.invocationKey(trimmed, name) + if pending := tracker.pendingByInvocation[key]; len(pending) > 0 { + publicID := pending[0] + if len(pending) == 1 { + delete(tracker.pendingByInvocation, key) + } else { + tracker.pendingByInvocation[key] = pending[1:] + } + return publicID + } + } + return tracker.nextCallID(trimmed, name) +} + +func (tracker *gigaChatResponsesCallIDTracker) nextCallID(toolsStateID string, name string) string { + ordinal := 0 + if tracker != nil { + ordinal = tracker.counts[toolsStateID] + tracker.counts[toolsStateID] = ordinal + 1 + } + return newGigaChatResponsesGeneratedCallID(toolsStateID, name, ordinal) +} + +func (tracker *gigaChatResponsesCallIDTracker) invocationKey(toolsStateID string, name string) string { + return toolsStateID + "\x00" + strings.TrimSpace(name) +} + +func newGigaChatResponsesGeneratedCallID(toolsStateID string, name string, ordinal int) string { + encodedToolsStateID := base64.RawURLEncoding.EncodeToString([]byte(toolsStateID)) + encodedName := base64.RawURLEncoding.EncodeToString([]byte(strings.TrimSpace(name))) + return fmt.Sprintf("%s%s.%s.%s.%d", gigaChatResponsesGeneratedCallIDPrefix, gigaChatResponsesGeneratedCallIDVersion, encodedToolsStateID, encodedName, ordinal) +} + +func toGigaChatResponsesMessageToolStateID(message GigaChatResponsesMessage) *string { + if message.ToolsStateID != nil && strings.TrimSpace(*message.ToolsStateID) != "" { + return schemas.Ptr(strings.TrimSpace(*message.ToolsStateID)) + } + if message.ToolStateID != nil && strings.TrimSpace(*message.ToolStateID) != "" { + return schemas.Ptr(strings.TrimSpace(*message.ToolStateID)) + } + return nil +} + +func toBifrostGigaChatResponsesItemID(prefix string, messageID *string, index int) string { + if messageID != nil && strings.TrimSpace(*messageID) != "" { + if index == 0 { + return strings.TrimSpace(*messageID) + } + return fmt.Sprintf("%s_%s_%d", prefix, strings.TrimSpace(*messageID), index) + } + return fmt.Sprintf("%s_%d", prefix, index) +} + +func toBifrostGigaChatResponsesFileItemID(prefix string, messageID *string, partIndex int, fileIndex int) string { + trimmedPrefix := strings.TrimSpace(prefix) + if trimmedPrefix == "" { + trimmedPrefix = "file" + } + if messageID != nil && strings.TrimSpace(*messageID) != "" { + if fileIndex == 0 { + return fmt.Sprintf("%s_%s_%d", trimmedPrefix, strings.TrimSpace(*messageID), partIndex) + } + return fmt.Sprintf("%s_%s_%d_%d", trimmedPrefix, strings.TrimSpace(*messageID), partIndex, fileIndex) + } + if fileIndex == 0 { + return fmt.Sprintf("%s_%d", trimmedPrefix, partIndex) + } + return fmt.Sprintf("%s_%d_%d", trimmedPrefix, partIndex, fileIndex) +} + +func toBifrostGigaChatResponsesReasoningItemID(messageID *string) *string { + if messageID == nil || strings.TrimSpace(*messageID) == "" { + return nil + } + itemID := "rs_" + strings.TrimSpace(*messageID) + return &itemID +} + +func isGigaChatResponsesReasoningRole(role string) bool { + return strings.EqualFold(strings.TrimSpace(role), gigaChatResponsesRoleReasoning) +} + +func stringifyGigaChatResponsesPayload(payload interface{}) string { + if payload == nil { + return "{}" + } + if text, ok := payload.(string); ok { + trimmed := strings.TrimSpace(text) + if trimmed == "" { + return "{}" + } + if sonic.ValidString(trimmed) { + return compactGigaChatResponsesJSON(trimmed) + } + return trimmed + } + + raw, err := sonic.ConfigStd.Marshal(payload) + if err != nil { + return "{}" + } + if sonic.Valid(raw) { + return compactGigaChatResponsesJSON(string(raw)) + } + return string(raw) +} + +func compactGigaChatResponsesJSON(raw string) string { + var builder strings.Builder + builder.Grow(len(raw)) + inString := false + escaped := false + for index := 0; index < len(raw); index++ { + character := raw[index] + if inString { + builder.WriteByte(character) + if escaped { + escaped = false + continue + } + if character == '\\' { + escaped = true + continue + } + if character == '"' { + inString = false + } + continue + } + switch character { + case '"': + inString = true + builder.WriteByte(character) + case ' ', '\n', '\r', '\t': + continue + default: + builder.WriteByte(character) + } + } + return builder.String() +} + +func toBifrostGigaChatResponsesStatus(finishReason *string) (*string, *schemas.ResponsesResponseIncompleteDetails, *string) { + mappedFinishReason := toBifrostGigaChatFinishReason(finishReason) + if mappedFinishReason == nil || strings.TrimSpace(*mappedFinishReason) == "" { + return nil, nil, nil + } + + stopReason := strings.TrimSpace(*mappedFinishReason) + switch stopReason { + case string(schemas.BifrostFinishReasonLength): + return schemas.Ptr("incomplete"), &schemas.ResponsesResponseIncompleteDetails{Reason: "max_output_tokens"}, &stopReason + default: + return schemas.Ptr("completed"), nil, &stopReason + } +} + +func toBifrostGigaChatResponsesProviderExtraFields(response *GigaChatResponsesResponse) map[string]interface{} { + if response == nil { + return nil + } + + fields := make(map[string]interface{}) + if response.ThreadID != nil && strings.TrimSpace(*response.ThreadID) != "" { + fields["thread_id"] = *response.ThreadID + } + if response.MessageID != nil && strings.TrimSpace(*response.MessageID) != "" { + fields["message_id"] = *response.MessageID + } + if response.ToolsStateID != nil && strings.TrimSpace(*response.ToolsStateID) != "" { + fields["tools_state_id"] = *response.ToolsStateID + } + if messageToolStateIDs := toBifrostGigaChatResponsesMessageToolStateIDs(response); len(messageToolStateIDs) > 0 { + fields["message_tools_state_ids"] = messageToolStateIDs + } + if response.ToolExecution != nil { + fields["tool_execution"] = response.ToolExecution + } + if response.AdditionalData != nil { + fields["additional_data"] = response.AdditionalData + } + if strings.TrimSpace(response.SystemFingerprint) != "" { + fields["system_fingerprint"] = response.SystemFingerprint + } + if len(response.ExtraParams) > 0 { + fields["gigachat_extra"] = response.ExtraParams + } + if len(fields) == 0 { + return nil + } + return fields +} + +func toBifrostGigaChatResponsesMessageToolStateIDs(response *GigaChatResponsesResponse) []map[string]interface{} { + if response == nil { + return nil + } + + if len(response.Messages) > 0 { + return collectBifrostGigaChatResponsesMessageToolStateIDs(response.Messages) + } + + if len(response.Choices) == 0 { + return nil + } + messages := make([]GigaChatResponsesMessage, 0, len(response.Choices)) + for _, choice := range response.Choices { + if choice.Message != nil { + messages = append(messages, *choice.Message) + } + } + return collectBifrostGigaChatResponsesMessageToolStateIDs(messages) +} + +func collectBifrostGigaChatResponsesMessageToolStateIDs(messages []GigaChatResponsesMessage) []map[string]interface{} { + messageToolStateIDs := make([]map[string]interface{}, 0) + for index, message := range messages { + toolsStateID := toGigaChatResponsesMessageToolStateID(message) + if toolsStateID == nil || strings.TrimSpace(*toolsStateID) == "" || hasGigaChatResponsesToolPayload(message) { + continue + } + + entry := map[string]interface{}{ + "index": index, + "tools_state_id": strings.TrimSpace(*toolsStateID), + } + if message.MessageID != nil && strings.TrimSpace(*message.MessageID) != "" { + entry["message_id"] = strings.TrimSpace(*message.MessageID) + } + if strings.TrimSpace(message.Role) != "" { + entry["role"] = strings.TrimSpace(message.Role) + } + messageToolStateIDs = append(messageToolStateIDs, entry) + } + return messageToolStateIDs +} + +func hasGigaChatResponsesToolPayload(message GigaChatResponsesMessage) bool { + if message.FunctionCall != nil { + return true + } + for _, part := range message.Content { + if part.FunctionCall != nil || part.FunctionResult != nil { + return true + } + } + return false +} + +func applyGigaChatResponsesParams(gigaChatReq *GigaChatResponsesRequest, params *schemas.ResponsesParameters) error { + modelOptions := &GigaChatResponsesModelOptions{ + Temperature: params.Temperature, + TopP: params.TopP, + MaxTokens: params.MaxOutputTokens, + TopLogProbs: params.TopLogProbs, + } + + if params.Reasoning != nil && params.Reasoning.Effort != nil && strings.TrimSpace(*params.Reasoning.Effort) != "" && *params.Reasoning.Effort != "none" { + modelOptions.Reasoning = &GigaChatResponsesReasoning{Effort: *params.Reasoning.Effort} + } + if params.Text != nil && params.Text.Format != nil { + responseFormat, err := toGigaChatResponsesResponseFormat(params.Text.Format) + if err != nil { + return err + } + modelOptions.ResponseFormat = responseFormat + } + if hasGigaChatResponsesModelOptions(modelOptions) { + gigaChatReq.ModelOptions = modelOptions + } + + toolsConversion, err := toGigaChatResponsesTools(params.Tools) + if err != nil { + return err + } + gigaChatReq.Tools = toolsConversion.Tools + if len(toolsConversion.UserInfo) > 0 { + gigaChatReq.UserInfo = toolsConversion.UserInfo + } + + toolConfig, err := toGigaChatResponsesToolConfig(params.ToolChoice, params.Tools) + if err != nil { + return err + } + gigaChatReq.ToolConfig = toolConfig + + if err := applyGigaChatResponsesStorage(gigaChatReq, params); err != nil { + return err + } + + return applyGigaChatResponsesExtraParams(gigaChatReq, params.ExtraParams) +} + +func applyGigaChatResponsesStorage(gigaChatReq *GigaChatResponsesRequest, params *schemas.ResponsesParameters) error { + if gigaChatReq == nil || params == nil { + return nil + } + + if params.Store != nil && !*params.Store { + if hasGigaChatResponsesStorageParams(params) { + return fmt.Errorf("GigaChat Responses cannot combine store=false with conversation, previous_response_id, or metadata") + } + gigaChatReq.Storage = false + return nil + } + + storage := &GigaChatResponsesStorage{} + conversationID := trimStringPtr(params.Conversation) + previousResponseID := trimStringPtr(params.PreviousResponseID) + switch { + case conversationID != "" && previousResponseID != "" && conversationID != previousResponseID: + return fmt.Errorf("GigaChat Responses requires conversation and previous_response_id to reference the same thread_id") + case conversationID != "": + storage.ThreadID = &conversationID + case previousResponseID != "": + storage.ThreadID = &previousResponseID + } + + if params.Metadata != nil && len(*params.Metadata) > 0 { + storage.Metadata = make(map[string]interface{}, len(*params.Metadata)) + for key, value := range *params.Metadata { + storage.Metadata[key] = value + } + } + + gigaChatReq.Storage = storage + return nil +} + +func hasGigaChatResponsesStorageParams(params *schemas.ResponsesParameters) bool { + if params == nil { + return false + } + return hasGigaChatResponsesThreadID(params) || + (params.Metadata != nil && len(*params.Metadata) > 0) +} + +func hasGigaChatResponsesThreadID(params *schemas.ResponsesParameters) bool { + if params == nil { + return false + } + return trimStringPtr(params.Conversation) != "" || + trimStringPtr(params.PreviousResponseID) != "" +} + +func trimStringPtr(value *string) string { + if value == nil { + return "" + } + return strings.TrimSpace(*value) +} + +func collectGigaChatResponsesFunctionCallNames(messages []schemas.ResponsesMessage) map[string]string { + functionNamesByID := make(map[string]string) + for _, message := range messages { + if message.ResponsesToolMessage == nil || message.ResponsesToolMessage.Name == nil { + continue + } + name := strings.TrimSpace(*message.ResponsesToolMessage.Name) + if name == "" { + continue + } + if message.ResponsesToolMessage.CallID != nil { + if callID := strings.TrimSpace(*message.ResponsesToolMessage.CallID); callID != "" { + functionNamesByID[callID] = name + if toolsStateID := toGigaChatResponsesToolsStateIDFromCallID(callID); toolsStateID != callID && toolsStateID != "" { + if _, exists := functionNamesByID[toolsStateID]; !exists { + functionNamesByID[toolsStateID] = name + } + } + } + } + if message.ID != nil { + if id := strings.TrimSpace(*message.ID); id != "" { + functionNamesByID[id] = name + } + } + } + return functionNamesByID +} + +func toGigaChatResponsesMessages(message schemas.ResponsesMessage, functionCallNamesByID map[string]string) ([]GigaChatResponsesMessage, error) { + messageType := schemas.ResponsesMessageTypeMessage + if message.Type != nil { + messageType = *message.Type + } + + switch messageType { + case schemas.ResponsesMessageTypeMessage: + return toGigaChatResponsesChatMessages(message) + case schemas.ResponsesMessageTypeFunctionCall: + return toGigaChatResponsesFunctionCallMessage(message) + case schemas.ResponsesMessageTypeFunctionCallOutput: + return toGigaChatResponsesFunctionResultMessage(message, functionCallNamesByID) + case schemas.ResponsesMessageTypeReasoning: + return toGigaChatResponsesReasoningMessage(message) + default: + return nil, fmt.Errorf("item type %q is not supported by GigaChat Responses", messageType) + } +} + +func toGigaChatResponsesChatMessages(message schemas.ResponsesMessage) ([]GigaChatResponsesMessage, error) { + role := schemas.ResponsesInputMessageRoleUser + if message.Role != nil { + role = *message.Role + } + switch role { + case schemas.ResponsesInputMessageRoleSystem, schemas.ResponsesInputMessageRoleUser, schemas.ResponsesInputMessageRoleAssistant: + case schemas.ResponsesInputMessageRoleDeveloper: + return nil, fmt.Errorf("developer messages are not supported by GigaChat Responses") + default: + return nil, fmt.Errorf("role %q is not supported by GigaChat Responses", role) + } + + content, err := toGigaChatResponsesContentParts(message.Content) + if err != nil { + return nil, err + } + return []GigaChatResponsesMessage{{ + Role: string(role), + MessageID: message.ID, + Content: content, + }}, nil +} + +func toGigaChatResponsesFunctionCallMessage(message schemas.ResponsesMessage) ([]GigaChatResponsesMessage, error) { + if message.ResponsesToolMessage == nil { + return nil, fmt.Errorf("function_call item requires tool message fields") + } + if message.ResponsesToolMessage.Name == nil || strings.TrimSpace(*message.ResponsesToolMessage.Name) == "" { + return nil, fmt.Errorf("function_call item name is required") + } + + arguments, err := parseGigaChatFunctionArguments(message.ResponsesToolMessage.Arguments) + if err != nil { + return nil, err + } + functionCall := &GigaChatResponsesFunctionCall{ + Name: toGigaChatResponsesFunctionName(*message.ResponsesToolMessage.Name), + Arguments: arguments, + } + return []GigaChatResponsesMessage{{ + Role: string(schemas.ResponsesInputMessageRoleAssistant), + MessageID: message.ID, + ToolsStateID: toGigaChatResponsesToolsStateID(message), + Content: []GigaChatResponsesContentPart{{ + FunctionCall: functionCall, + }}, + FunctionCall: functionCall, + }}, nil +} + +func toGigaChatResponsesFunctionResultMessage(message schemas.ResponsesMessage, functionCallNamesByID map[string]string) ([]GigaChatResponsesMessage, error) { + if message.ResponsesToolMessage == nil { + return nil, fmt.Errorf("function_call_output item requires tool message fields") + } + name := "" + if message.ResponsesToolMessage.Name != nil { + name = strings.TrimSpace(*message.ResponsesToolMessage.Name) + } + if name == "" { + name = functionCallNamesByID[trimStringPtr(message.ResponsesToolMessage.CallID)] + } + if name == "" { + name = functionCallNamesByID[trimStringPtr(message.ID)] + } + if name == "" { + name = toGigaChatResponsesFunctionNameFromCallID(trimStringPtr(message.ResponsesToolMessage.CallID)) + } + if name == "" { + return nil, fmt.Errorf("function_call_output item name is required") + } + + result, err := toGigaChatFunctionResultPayload(message) + if err != nil { + return nil, err + } + return []GigaChatResponsesMessage{{ + Role: "tool", + MessageID: message.ID, + ToolsStateID: toGigaChatResponsesToolsStateID(message), + Content: []GigaChatResponsesContentPart{{ + FunctionResult: &GigaChatResponsesFunctionResult{ + Name: toGigaChatResponsesFunctionName(name), + Result: result, + }, + }}, + }}, nil +} + +func toGigaChatResponsesToolsStateID(message schemas.ResponsesMessage) *string { + if message.ResponsesToolMessage != nil && message.ResponsesToolMessage.CallID != nil && strings.TrimSpace(*message.ResponsesToolMessage.CallID) != "" { + return schemas.Ptr(toGigaChatResponsesToolsStateIDFromCallID(*message.ResponsesToolMessage.CallID)) + } + if message.ID != nil && strings.TrimSpace(*message.ID) != "" { + return schemas.Ptr(strings.TrimSpace(*message.ID)) + } + return nil +} + +func toGigaChatResponsesToolsStateIDFromCallID(callID string) string { + trimmed := strings.TrimSpace(callID) + if trimmed == "" { + return "" + } + toolsStateID, _, ok := decodeGigaChatResponsesGeneratedCallID(trimmed) + if !ok { + return trimmed + } + return toolsStateID +} + +func toGigaChatResponsesFunctionNameFromCallID(callID string) string { + _, name, ok := decodeGigaChatResponsesGeneratedCallID(callID) + if !ok { + return "" + } + return name +} + +func decodeGigaChatResponsesGeneratedCallID(callID string) (string, string, bool) { + trimmed := strings.TrimSpace(callID) + if trimmed == "" || !strings.HasPrefix(trimmed, gigaChatResponsesGeneratedCallIDPrefix) { + return "", "", false + } + + encoded := strings.TrimPrefix(trimmed, gigaChatResponsesGeneratedCallIDPrefix) + parts := strings.Split(encoded, ".") + if (len(parts) != 3 && len(parts) != 4) || parts[0] != gigaChatResponsesGeneratedCallIDVersion { + return "", "", false + } + ordinalIndex := len(parts) - 1 + if _, err := strconv.Atoi(parts[ordinalIndex]); err != nil { + return "", "", false + } + decodedToolsStateID, err := base64.RawURLEncoding.DecodeString(parts[1]) + if err != nil { + return "", "", false + } + toolsStateID := strings.TrimSpace(string(decodedToolsStateID)) + if toolsStateID == "" { + return "", "", false + } + if len(parts) == 3 { + return toolsStateID, "", true + } + decodedName, err := base64.RawURLEncoding.DecodeString(parts[2]) + if err != nil { + return "", "", false + } + return toolsStateID, strings.TrimSpace(string(decodedName)), true +} + +func toGigaChatResponsesReasoningMessage(message schemas.ResponsesMessage) ([]GigaChatResponsesMessage, error) { + content := make([]GigaChatResponsesContentPart, 0) + if message.ResponsesReasoning != nil { + for _, summary := range message.ResponsesReasoning.Summary { + text := summary.Text + content = append(content, GigaChatResponsesContentPart{Text: &text}) + } + } + if message.Content != nil { + parts, err := toGigaChatResponsesContentParts(message.Content) + if err != nil { + return nil, err + } + content = append(content, parts...) + } + if len(content) == 0 { + return nil, fmt.Errorf("reasoning item content is required") + } + return []GigaChatResponsesMessage{{ + Role: gigaChatResponsesRoleReasoning, + MessageID: message.ID, + Content: content, + }}, nil +} + +func toGigaChatResponsesContentParts(content *schemas.ResponsesMessageContent) ([]GigaChatResponsesContentPart, error) { + if content == nil { + return nil, nil + } + if content.ContentStr != nil { + return []GigaChatResponsesContentPart{{Text: content.ContentStr}}, nil + } + if content.ContentBlocks == nil { + return nil, nil + } + + parts := make([]GigaChatResponsesContentPart, 0, len(content.ContentBlocks)) + for index, block := range content.ContentBlocks { + switch block.Type { + case schemas.ResponsesInputMessageContentBlockTypeText, + schemas.ResponsesOutputMessageContentTypeText, + schemas.ResponsesOutputMessageContentTypeReasoning: + if block.Text != nil { + parts = append(parts, GigaChatResponsesContentPart{Text: block.Text}) + } + case schemas.ResponsesInputMessageContentBlockTypeFile: + file, err := toGigaChatResponsesContentFile(index, block) + if err != nil { + return nil, err + } + parts = append(parts, GigaChatResponsesContentPart{Files: []GigaChatResponsesContentFile{*file}}) + case schemas.ResponsesInputMessageContentBlockTypeImage: + file, err := toGigaChatResponsesContentImage(index, block) + if err != nil { + return nil, err + } + parts = append(parts, GigaChatResponsesContentPart{Files: []GigaChatResponsesContentFile{*file}}) + case schemas.ResponsesInputMessageContentBlockTypeAudio: + return nil, fmt.Errorf("content block %d: input_audio is not supported by GigaChat Responses request conversion yet", index) + default: + return nil, fmt.Errorf("content block %d: type %q is not supported by GigaChat Responses", index, block.Type) + } + } + return parts, nil +} + +func toGigaChatResponsesContentImage(index int, block schemas.ResponsesMessageContentBlock) (*GigaChatResponsesContentFile, error) { + if block.ResponsesInputMessageContentBlockImage != nil && + block.ResponsesInputMessageContentBlockImage.ImageURL != nil && + strings.TrimSpace(*block.ResponsesInputMessageContentBlockImage.ImageURL) != "" { + return nil, fmt.Errorf("content block %d: GigaChat Responses supports pre-uploaded file_id references for input_image; upload image_url with the Files API before calling Responses", index) + } + + fileID := "" + if block.FileID != nil { + fileID = strings.TrimSpace(*block.FileID) + } + if fileID == "" { + return nil, fmt.Errorf("content block %d: input_image requires file_id; upload image_url with the Files API before calling Responses", index) + } + return &GigaChatResponsesContentFile{ID: fileID}, nil +} + +func toGigaChatResponsesContentFile(index int, block schemas.ResponsesMessageContentBlock) (*GigaChatResponsesContentFile, error) { + if block.ResponsesInputMessageContentBlockFile != nil && (block.FileData != nil || block.FileURL != nil) { + return nil, fmt.Errorf("content block %d: GigaChat Responses supports pre-uploaded file_id references only; upload inline file content with the Files API before calling Responses", index) + } + + fileID := "" + if block.FileID != nil { + fileID = strings.TrimSpace(*block.FileID) + } + if fileID == "" { + return nil, fmt.Errorf("content block %d: GigaChat file content requires file_id; upload the file with the Files API before calling Responses", index) + } + + file := &GigaChatResponsesContentFile{ID: fileID} + if block.ResponsesInputMessageContentBlockFile != nil && block.FileType != nil { + if mime := strings.TrimSpace(*block.FileType); mime != "" { + file.MIME = &mime + } + } + return file, nil +} + +func parseGigaChatFunctionArguments(arguments *string) (interface{}, error) { + if arguments == nil || strings.TrimSpace(*arguments) == "" { + return map[string]interface{}{}, nil + } + var parsed map[string]interface{} + if err := json.Unmarshal([]byte(*arguments), &parsed); err != nil { + return nil, fmt.Errorf("function_call arguments must be a JSON object: %w", err) + } + return parsed, nil +} + +func toGigaChatFunctionResultPayload(message schemas.ResponsesMessage) (interface{}, error) { + if message.ResponsesToolMessage != nil && message.ResponsesToolMessage.Output != nil { + output := message.ResponsesToolMessage.Output + if output.ResponsesToolCallOutputStr != nil { + return *output.ResponsesToolCallOutputStr, nil + } + if output.ResponsesFunctionToolCallOutputBlocks != nil { + return textFromGigaChatResponsesBlocks(output.ResponsesFunctionToolCallOutputBlocks) + } + } + if message.Content != nil { + if message.Content.ContentStr != nil { + return *message.Content.ContentStr, nil + } + if message.Content.ContentBlocks != nil { + return textFromGigaChatResponsesBlocks(message.Content.ContentBlocks) + } + } + return "", nil +} + +func textFromGigaChatResponsesBlocks(blocks []schemas.ResponsesMessageContentBlock) (string, error) { + var builder strings.Builder + for index, block := range blocks { + switch block.Type { + case schemas.ResponsesInputMessageContentBlockTypeText, schemas.ResponsesOutputMessageContentTypeText: + if block.Text != nil { + builder.WriteString(*block.Text) + } + default: + return "", fmt.Errorf("function result block %d with type %q is not supported by GigaChat Responses", index, block.Type) + } + } + return builder.String(), nil +} + +func toGigaChatResponsesResponseFormat(format *schemas.ResponsesTextConfigFormat) (*GigaChatResponsesResponseFormat, error) { + switch format.Type { + case "text": + return &GigaChatResponsesResponseFormat{Type: "text"}, nil + case "json_schema": + if format.JSONSchema == nil { + return nil, fmt.Errorf("response_format json_schema requires schema") + } + schema, err := cloneGigaChatSchemaValue(format.JSONSchema.ToMap()) + if err != nil { + return nil, fmt.Errorf("response_format json_schema is invalid: %w", err) + } + if schema == nil { + return nil, fmt.Errorf("response_format json_schema requires non-empty schema") + } + schema = withGigaChatResponseFormatSchemaMetadata(schema, format.Name, format.Description) + strict := format.Strict + if strict == nil && format.JSONSchema.Strict != nil { + strict = format.JSONSchema.Strict + } + return &GigaChatResponsesResponseFormat{ + Type: "json_schema", + Schema: schema, + Strict: strict, + }, nil + default: + return nil, fmt.Errorf("response_format type %q is not supported by GigaChat Responses", format.Type) + } +} + +func withGigaChatResponseFormatSchemaMetadata(schema interface{}, name *string, description *string) interface{} { + switch schemaMap := schema.(type) { + case map[string]interface{}: + if name != nil && strings.TrimSpace(*name) != "" { + if _, exists := schemaMap["title"]; !exists { + schemaMap["title"] = strings.TrimSpace(*name) + } + } + if description != nil && strings.TrimSpace(*description) != "" { + if _, exists := schemaMap["description"]; !exists { + schemaMap["description"] = strings.TrimSpace(*description) + } + } + return schemaMap + case *schemas.OrderedMap: + if schemaMap == nil { + return schema + } + if name != nil && strings.TrimSpace(*name) != "" { + if _, exists := schemaMap.Get("title"); !exists { + schemaMap.Set("title", strings.TrimSpace(*name)) + } + } + if description != nil && strings.TrimSpace(*description) != "" { + if _, exists := schemaMap.Get("description"); !exists { + schemaMap.Set("description", strings.TrimSpace(*description)) + } + } + return schemaMap + case schemas.OrderedMap: + schemaCopy := schemaMap.Clone() + return withGigaChatResponseFormatSchemaMetadata(schemaCopy, name, description) + default: + return schema + } +} + +func applyGigaChatResponsesExtraParams(gigaChatReq *GigaChatResponsesRequest, extraParams map[string]interface{}) error { + if len(extraParams) == 0 { + return nil + } + + remaining := make(map[string]interface{}, len(extraParams)) + for name, value := range extraParams { + remaining[name] = value + } + + if value, ok, err := consumeStringExtraParam(remaining, "assistant_id"); err != nil { + return err + } else if ok { + gigaChatReq.AssistantID = &value + } + if value, ok, err := consumeStringExtraParam(remaining, "tools_state_id"); err != nil { + return err + } else if ok { + gigaChatReq.ToolsStateID = &value + } + if value, ok, err := consumeBoolExtraParam(remaining, "disable_filter"); err != nil { + return err + } else if ok { + gigaChatReq.DisableFilter = &value + } + if value, ok, err := consumeStringSliceExtraParam(remaining, "flags"); err != nil { + return err + } else if ok { + gigaChatReq.Flags = value + } + if value, ok, err := consumeMapExtraParam(remaining, "filter_config"); err != nil { + return err + } else if ok { + gigaChatReq.FilterConfig = value + } + if value, ok, err := consumeMapExtraParam(remaining, "ranker_options"); err != nil { + return err + } else if ok { + gigaChatReq.RankerOptions = value + } + if value, ok, err := consumeMapExtraParam(remaining, "user_info"); err != nil { + return err + } else if ok { + gigaChatReq.UserInfo = value + } + if value, ok := remaining["storage"]; ok { + gigaChatReq.Storage = value + delete(remaining, "storage") + } + + modelOptions := ensureGigaChatResponsesModelOptions(gigaChatReq) + if value, ok, err := consumeStringExtraParam(remaining, "preset"); err != nil { + return err + } else if ok { + modelOptions.Preset = &value + } + if value, ok, err := consumeFloatExtraParam(remaining, "repetition_penalty"); err != nil { + return err + } else if ok { + modelOptions.RepetitionPenalty = &value + } + if value, ok, err := consumeFloatExtraParam(remaining, "update_interval"); err != nil { + return err + } else if ok { + modelOptions.UpdateInterval = &value + } + if value, ok, err := consumeBoolExtraParam(remaining, "unnormalized_history"); err != nil { + return err + } else if ok { + modelOptions.UnnormalizedHistory = &value + } + if !hasGigaChatResponsesModelOptions(modelOptions) { + gigaChatReq.ModelOptions = nil + } + if len(remaining) > 0 { + gigaChatReq.ExtraParams = remaining + } + return nil +} + +func unsupportedGigaChatResponsesParams(params *schemas.ResponsesParameters) []string { + if params == nil { + return nil + } + + unsupported := make([]string, 0) + addIf := func(condition bool, name string) { + if condition { + unsupported = append(unsupported, name) + } + } + + addIf(params.Background != nil, "background") + addIf(len(params.Include) > 0, "include") + addIf(params.MaxToolCalls != nil, "max_tool_calls") + addIf(params.ParallelToolCalls != nil && *params.ParallelToolCalls, "parallel_tool_calls") + addIf(params.PromptCacheKey != nil, "prompt_cache_key") + addIf(params.SafetyIdentifier != nil, "safety_identifier") + addIf(params.ServiceTier != nil, "service_tier") + addIf(params.StreamOptions != nil, "stream_options") + addIf(params.Truncation != nil, "truncation") + addIf(params.User != nil, "user") + if params.Reasoning != nil { + addIf(params.Reasoning.GenerateSummary != nil, "reasoning.generate_summary") + addIf(params.Reasoning.Summary != nil, "reasoning.summary") + addIf(params.Reasoning.MaxTokens != nil, "reasoning.max_tokens") + } + if params.Text != nil { + addIf(params.Text.Verbosity != nil, "text.verbosity") + } + unsupported = append(unsupported, unsupportedGigaChatToolControlExtraParams(params.ExtraParams, "functions", "function_call", "tools", "tool_config", "parallel_tool_calls")...) + + sort.Strings(unsupported) + return unsupported +} + +func ensureGigaChatResponsesModelOptions(gigaChatReq *GigaChatResponsesRequest) *GigaChatResponsesModelOptions { + if gigaChatReq.ModelOptions == nil { + gigaChatReq.ModelOptions = &GigaChatResponsesModelOptions{} + } + return gigaChatReq.ModelOptions +} + +func hasGigaChatResponsesModelOptions(options *GigaChatResponsesModelOptions) bool { + if options == nil { + return false + } + return options.Preset != nil || + options.Temperature != nil || + options.TopP != nil || + options.MaxTokens != nil || + options.RepetitionPenalty != nil || + options.UpdateInterval != nil || + options.UnnormalizedHistory != nil || + options.TopLogProbs != nil || + options.Reasoning != nil || + options.ResponseFormat != nil || + len(options.ExtraParams) > 0 +} + +func consumeStringExtraParam(params map[string]interface{}, name string) (string, bool, error) { + value, ok := params[name] + if !ok { + return "", false, nil + } + converted, ok := schemas.SafeExtractString(value) + if !ok { + return "", true, fmt.Errorf("extra parameter %q must be a string", name) + } + delete(params, name) + return strings.TrimSpace(converted), true, nil +} + +func consumeBoolExtraParam(params map[string]interface{}, name string) (bool, bool, error) { + value, ok := params[name] + if !ok { + return false, false, nil + } + converted, ok := schemas.SafeExtractBool(value) + if !ok { + return false, true, fmt.Errorf("extra parameter %q must be a boolean", name) + } + delete(params, name) + return converted, true, nil +} + +func consumeFloatExtraParam(params map[string]interface{}, name string) (float64, bool, error) { + value, ok := params[name] + if !ok { + return 0, false, nil + } + converted, ok := schemas.SafeExtractFloat64(value) + if !ok { + return 0, true, fmt.Errorf("extra parameter %q must be a number", name) + } + delete(params, name) + return converted, true, nil +} + +func consumeStringSliceExtraParam(params map[string]interface{}, name string) ([]string, bool, error) { + value, ok := params[name] + if !ok { + return nil, false, nil + } + converted, ok := schemas.SafeExtractStringSlice(value) + if !ok { + return nil, true, fmt.Errorf("extra parameter %q must be an array of strings", name) + } + delete(params, name) + return converted, true, nil +} + +func consumeMapExtraParam(params map[string]interface{}, name string) (map[string]interface{}, bool, error) { + value, ok := params[name] + if !ok { + return nil, false, nil + } + converted, ok := value.(map[string]interface{}) + if !ok { + return nil, true, fmt.Errorf("extra parameter %q must be an object", name) + } + delete(params, name) + return converted, true, nil +} diff --git a/core/providers/gigachat/responses_attachments.go b/core/providers/gigachat/responses_attachments.go new file mode 100644 index 00000000000..9dbbbc557b4 --- /dev/null +++ b/core/providers/gigachat/responses_attachments.go @@ -0,0 +1,304 @@ +package gigachat + +import ( + "context" + "fmt" + "mime" + "net/url" + "path" + "path/filepath" + "strings" + + providerUtils "github.com/maximhq/bifrost/core/providers/utils" + "github.com/maximhq/bifrost/core/schemas" +) + +type gigaChatResourceFetchFunc func(context.Context, string) (string, string, error) + +func (provider *GigaChatProvider) prepareGigaChatResponsesAttachments(ctx *schemas.BifrostContext, key schemas.Key, request *schemas.BifrostResponsesRequest) (*schemas.BifrostResponsesRequest, *schemas.BifrostError) { + if request == nil { + return nil, providerUtils.NewBifrostOperationError("responses request is nil", nil) + } + + var prepared *schemas.BifrostResponsesRequest + for messageIndex := range request.Input { + content := request.Input[messageIndex].Content + if content == nil || len(content.ContentBlocks) == 0 { + continue + } + + for blockIndex, block := range content.ContentBlocks { + if gigaChatResponsesAttachmentMayUpload(block) { + if replacement, ok := provider.getCachedGigaChatResponsesAttachment(ctx, key, request, messageIndex, blockIndex); ok { + if prepared == nil { + prepared = cloneGigaChatResponsesRequestForAttachmentUpload(request) + } + prepared.Input[messageIndex].Content.ContentBlocks[blockIndex] = replacement + continue + } + } + + replacement, changed, bifrostErr := provider.prepareGigaChatResponsesAttachmentBlock(ctx, key, blockIndex, block) + if bifrostErr != nil { + return nil, bifrostErr + } + if !changed { + continue + } + + if prepared == nil { + prepared = cloneGigaChatResponsesRequestForAttachmentUpload(request) + } + provider.setCachedGigaChatResponsesAttachment(ctx, key, request, messageIndex, blockIndex, replacement) + prepared.Input[messageIndex].Content.ContentBlocks[blockIndex] = replacement + } + } + + if prepared != nil { + return prepared, nil + } + return request, nil +} + +func gigaChatResponsesAttachmentMayUpload(block schemas.ResponsesMessageContentBlock) bool { + switch block.Type { + case schemas.ResponsesInputMessageContentBlockTypeImage: + return block.ResponsesInputMessageContentBlockImage != nil && + block.ResponsesInputMessageContentBlockImage.ImageURL != nil && + strings.TrimSpace(*block.ResponsesInputMessageContentBlockImage.ImageURL) != "" + case schemas.ResponsesInputMessageContentBlockTypeFile: + file := block.ResponsesInputMessageContentBlockFile + return file != nil && + ((file.FileData != nil && strings.TrimSpace(*file.FileData) != "") || + (file.FileURL != nil && strings.TrimSpace(*file.FileURL) != "")) + default: + return false + } +} + +func cloneGigaChatResponsesRequestForAttachmentUpload(request *schemas.BifrostResponsesRequest) *schemas.BifrostResponsesRequest { + prepared := *request + prepared.Input = make([]schemas.ResponsesMessage, len(request.Input)) + for i := range request.Input { + prepared.Input[i] = schemas.DeepCopyResponsesMessage(request.Input[i]) + } + return &prepared +} + +func (provider *GigaChatProvider) prepareGigaChatResponsesAttachmentBlock(ctx *schemas.BifrostContext, key schemas.Key, blockIndex int, block schemas.ResponsesMessageContentBlock) (schemas.ResponsesMessageContentBlock, bool, *schemas.BifrostError) { + switch block.Type { + case schemas.ResponsesInputMessageContentBlockTypeImage: + if block.ResponsesInputMessageContentBlockImage == nil || + block.ResponsesInputMessageContentBlockImage.ImageURL == nil || + strings.TrimSpace(*block.ResponsesInputMessageContentBlockImage.ImageURL) == "" { + return block, false, nil + } + + upload, err := gigaChatResponsesImageURLUpload(ctx, blockIndex, block, providerUtils.FetchAndEncodeURL) + if err != nil { + return schemas.ResponsesMessageContentBlock{}, false, providerUtils.NewBifrostOperationError(err.Error(), err) + } + return provider.uploadGigaChatResponsesAttachment(ctx, key, upload) + case schemas.ResponsesInputMessageContentBlockTypeFile: + file := block.ResponsesInputMessageContentBlockFile + if file == nil { + return block, false, nil + } + switch { + case file.FileData != nil && strings.TrimSpace(*file.FileData) != "": + upload, err := gigaChatResponsesInlineFileUpload(blockIndex, file) + if err != nil { + return schemas.ResponsesMessageContentBlock{}, false, providerUtils.NewBifrostOperationError(err.Error(), err) + } + return provider.uploadGigaChatResponsesAttachment(ctx, key, upload) + case file.FileURL != nil && strings.TrimSpace(*file.FileURL) != "": + upload, err := gigaChatResponsesFileURLUpload(ctx, blockIndex, file, providerUtils.FetchAndEncodeURL) + if err != nil { + return schemas.ResponsesMessageContentBlock{}, false, providerUtils.NewBifrostOperationError(err.Error(), err) + } + return provider.uploadGigaChatResponsesAttachment(ctx, key, upload) + default: + return block, false, nil + } + default: + return block, false, nil + } +} + +func (provider *GigaChatProvider) uploadGigaChatResponsesAttachment(ctx *schemas.BifrostContext, key schemas.Key, upload gigaChatChatAttachmentUpload) (schemas.ResponsesMessageContentBlock, bool, *schemas.BifrostError) { + uploadResp, bifrostErr := provider.FileUpload(ctx, key, &schemas.BifrostFileUploadRequest{ + Provider: provider.GetProviderKey(), + File: upload.file, + Filename: upload.filename, + Purpose: schemas.FilePurposeUserData, + ContentType: &upload.contentType, + }) + if bifrostErr != nil { + return schemas.ResponsesMessageContentBlock{}, false, bifrostErr + } + if uploadResp == nil || strings.TrimSpace(uploadResp.ID) == "" { + return schemas.ResponsesMessageContentBlock{}, false, providerUtils.NewBifrostOperationError("GigaChat file upload response did not include file id", nil) + } + + fileID := strings.TrimSpace(uploadResp.ID) + return gigaChatResponsesUploadedFileBlock(fileID, upload.filename, upload.contentType), true, nil +} + +func gigaChatResponsesUploadedFileBlock(fileID string, filename string, contentType string) schemas.ResponsesMessageContentBlock { + file := &schemas.ResponsesInputMessageContentBlockFile{} + if trimmedFilename := strings.TrimSpace(filename); trimmedFilename != "" { + file.Filename = &trimmedFilename + } + if trimmedContentType := strings.TrimSpace(contentType); trimmedContentType != "" { + file.FileType = &trimmedContentType + } + return schemas.ResponsesMessageContentBlock{ + Type: schemas.ResponsesInputMessageContentBlockTypeFile, + FileID: schemas.Ptr(strings.TrimSpace(fileID)), + ResponsesInputMessageContentBlockFile: file, + } +} + +func gigaChatResponsesImageURLUpload(ctx context.Context, blockIndex int, block schemas.ResponsesMessageContentBlock, fetch gigaChatResourceFetchFunc) (gigaChatChatAttachmentUpload, error) { + if block.ResponsesInputMessageContentBlockImage == nil || block.ResponsesInputMessageContentBlockImage.ImageURL == nil { + return gigaChatChatAttachmentUpload{}, fmt.Errorf("content block %d: image_url is required", blockIndex) + } + + imageURL := strings.TrimSpace(*block.ResponsesInputMessageContentBlockImage.ImageURL) + sanitizedURL, err := schemas.SanitizeImageURL(imageURL) + if err != nil { + return gigaChatChatAttachmentUpload{}, fmt.Errorf("content block %d: invalid image_url: %w", blockIndex, err) + } + + urlInfo := schemas.ExtractURLTypeInfo(sanitizedURL) + contentType := "image/jpeg" + if urlInfo.MediaType != nil && strings.TrimSpace(*urlInfo.MediaType) != "" { + contentType = normalizeGigaChatContentType(*urlInfo.MediaType) + } + + if urlInfo.Type == schemas.ImageContentTypeBase64 { + if urlInfo.DataURLWithoutPrefix == nil { + return gigaChatChatAttachmentUpload{}, fmt.Errorf("content block %d: image_url base64 payload is required", blockIndex) + } + fileBytes, err := decodeGigaChatAttachmentBase64(*urlInfo.DataURLWithoutPrefix) + if err != nil { + return gigaChatChatAttachmentUpload{}, fmt.Errorf("content block %d: failed to decode image_url data: %w", blockIndex, err) + } + return gigaChatChatAttachmentUpload{ + file: fileBytes, + filename: "image" + extensionForGigaChatContentType(contentType), + contentType: contentType, + }, nil + } + + if fetch == nil { + fetch = providerUtils.FetchAndEncodeURL + } + fetchedContentType, fetchedBase64, err := fetch(ctx, sanitizedURL) + if err != nil { + return gigaChatChatAttachmentUpload{}, fmt.Errorf("content block %d: failed to fetch image_url: %w", blockIndex, err) + } + if fetchedContentType = normalizeGigaChatFetchedContentType(fetchedContentType); fetchedContentType != "" { + contentType = fetchedContentType + } + + fileBytes, err := decodeGigaChatAttachmentBase64(fetchedBase64) + if err != nil { + return gigaChatChatAttachmentUpload{}, fmt.Errorf("content block %d: failed to decode fetched image_url data: %w", blockIndex, err) + } + return gigaChatChatAttachmentUpload{ + file: fileBytes, + filename: filenameForGigaChatRemoteAttachment(sanitizedURL, "image", contentType), + contentType: contentType, + }, nil +} + +func gigaChatResponsesInlineFileUpload(blockIndex int, file *schemas.ResponsesInputMessageContentBlockFile) (gigaChatChatAttachmentUpload, error) { + if file == nil { + return gigaChatChatAttachmentUpload{}, fmt.Errorf("content block %d: file block is missing file payload", blockIndex) + } + return gigaChatChatFileUpload(blockIndex, &schemas.ChatInputFile{ + Filename: file.Filename, + FileData: file.FileData, + FileType: file.FileType, + }) +} + +func gigaChatResponsesFileURLUpload(ctx context.Context, blockIndex int, file *schemas.ResponsesInputMessageContentBlockFile, fetch gigaChatResourceFetchFunc) (gigaChatChatAttachmentUpload, error) { + if file == nil || file.FileURL == nil || strings.TrimSpace(*file.FileURL) == "" { + return gigaChatChatAttachmentUpload{}, fmt.Errorf("content block %d: file_url is required", blockIndex) + } + + fileURL := strings.TrimSpace(*file.FileURL) + if fetch == nil { + fetch = providerUtils.FetchAndEncodeURL + } + fetchedContentType, fetchedBase64, err := fetch(ctx, fileURL) + if err != nil { + return gigaChatChatAttachmentUpload{}, fmt.Errorf("content block %d: failed to fetch file_url: %w", blockIndex, err) + } + + fileBytes, err := decodeGigaChatAttachmentBase64(fetchedBase64) + if err != nil { + return gigaChatChatAttachmentUpload{}, fmt.Errorf("content block %d: failed to decode fetched file_url data: %w", blockIndex, err) + } + + filename := strings.TrimSpace(valueOrEmpty(file.Filename)) + if filename == "" { + filename = filenameFromGigaChatRemoteURL(fileURL) + } + contentType := strings.TrimSpace(valueOrEmpty(file.FileType)) + if contentType == "" { + contentType = normalizeGigaChatFetchedContentType(fetchedContentType) + } + if contentType == "" { + contentType = inferGigaChatContentTypeFromName(filename) + } + contentType = normalizeGigaChatContentType(contentType) + + return gigaChatChatAttachmentUpload{ + file: fileBytes, + filename: filenameForGigaChatAttachment(filename, contentType, "file"), + contentType: contentType, + }, nil +} + +func normalizeGigaChatFetchedContentType(contentType string) string { + contentType = strings.TrimSpace(contentType) + if contentType == "" { + return "" + } + if mediaType, _, err := mime.ParseMediaType(contentType); err == nil { + contentType = mediaType + } + return normalizeGigaChatContentType(contentType) +} + +func filenameForGigaChatRemoteAttachment(resourceURL string, fallbackBase string, contentType string) string { + if filename := filenameFromGigaChatRemoteURL(resourceURL); filename != "" { + return filename + } + return filenameForGigaChatAttachment("", contentType, fallbackBase) +} + +func filenameFromGigaChatRemoteURL(resourceURL string) string { + parsedURL, err := url.Parse(resourceURL) + if err != nil { + return "" + } + filename := path.Base(parsedURL.Path) + if filename == "." || filename == "/" { + return "" + } + return strings.TrimSpace(filename) +} + +func inferGigaChatContentTypeFromName(filename string) string { + if filename == "" { + return "" + } + if contentType := mime.TypeByExtension(strings.ToLower(filepath.Ext(filename))); contentType != "" { + return normalizeGigaChatFetchedContentType(contentType) + } + return normalizeGigaChatFetchedContentType(mime.TypeByExtension(strings.ToLower(path.Ext(filename)))) +} diff --git a/core/providers/gigachat/schema.go b/core/providers/gigachat/schema.go new file mode 100644 index 00000000000..a6cfc1a8bda --- /dev/null +++ b/core/providers/gigachat/schema.go @@ -0,0 +1,444 @@ +package gigachat + +import ( + "fmt" + "strings" + + schemas "github.com/maximhq/bifrost/core/schemas" +) + +func sanitizeGigaChatFunctionSchema(schema interface{}) (*schemas.ToolFunctionParameters, error) { + if schema == nil { + return nil, fmt.Errorf("function parameters JSON schema is required") + } + + raw, err := schemas.MarshalSorted(schema) + if err != nil { + return nil, fmt.Errorf("function parameters JSON schema is invalid: %w", err) + } + + var root schemas.OrderedMap + if err := schemas.Unmarshal(raw, &root); err != nil { + return nil, fmt.Errorf("function parameters JSON schema is invalid: %w", err) + } + + sanitized, nullOnly, err := sanitizeGigaChatSchemaMap(&root, &root, make(map[string]bool), "$") + if err != nil { + return nil, err + } + if nullOnly { + return nil, fmt.Errorf("function parameters JSON schema cannot be null-only") + } + + sanitizedRaw, err := schemas.MarshalSorted(sanitized) + if err != nil { + return nil, fmt.Errorf("function parameters JSON schema is invalid: %w", err) + } + + var parameters schemas.ToolFunctionParameters + if err := schemas.Unmarshal(sanitizedRaw, ¶meters); err != nil { + return nil, fmt.Errorf("function parameters JSON schema is invalid after GigaChat sanitization: %w", err) + } + return ¶meters, nil +} + +func sanitizeGigaChatSchemaMap(schema *schemas.OrderedMap, root *schemas.OrderedMap, resolving map[string]bool, path string) (*schemas.OrderedMap, bool, error) { + if schema == nil { + return nil, false, fmt.Errorf("%s: JSON schema object is required", path) + } + + if refValue, ok := schema.Get("$ref"); ok { + ref, ok := refValue.(string) + if !ok || strings.TrimSpace(ref) == "" { + return nil, false, fmt.Errorf("%s: $ref must be a non-empty string", path) + } + ref = strings.TrimSpace(ref) + if resolving[ref] { + return nil, false, fmt.Errorf("%s: circular local $ref %q is not supported", path, ref) + } + resolved, ok := lookupGigaChatLocalSchemaRef(root, ref) + if !ok { + return nil, false, fmt.Errorf("%s: unsupported or unresolved local $ref %q", path, ref) + } + resolving[ref] = true + merged, err := cloneGigaChatSchemaMap(resolved) + if err != nil { + delete(resolving, ref) + return nil, false, fmt.Errorf("%s: clone $ref %q: %w", path, ref, err) + } + var copyErr error + schema.Range(func(key string, value interface{}) bool { + if key == "$ref" || key == "$defs" || key == "definitions" { + return true + } + copied, err := cloneGigaChatSchemaValue(value) + if err != nil { + copyErr = err + return false + } + merged.Set(key, copied) + return true + }) + if copyErr != nil { + delete(resolving, ref) + return nil, false, fmt.Errorf("%s: merge $ref siblings: %w", path, copyErr) + } + out, nullOnly, err := sanitizeGigaChatSchemaMap(merged, root, resolving, path) + delete(resolving, ref) + return out, nullOnly, err + } + + for _, keyword := range []string{"anyOf", "oneOf"} { + if _, ok := schema.Get(keyword); ok { + return sanitizeGigaChatSchemaComposition(schema, root, resolving, path, keyword) + } + } + + if nullOnly, err := sanitizeGigaChatSchemaType(schema, path); err != nil || nullOnly { + return nil, nullOnly, err + } + + schema.Delete("nullable") + + if err := sanitizeGigaChatSchemaProperties(schema, root, resolving, path); err != nil { + return nil, false, err + } + if err := sanitizeGigaChatSchemaItems(schema, root, resolving, path); err != nil { + return nil, false, err + } + if err := sanitizeGigaChatSchemaAdditionalProperties(schema, root, resolving, path); err != nil { + return nil, false, err + } + if nullOnly, err := sanitizeGigaChatSchemaAllOf(schema, root, resolving, path); err != nil || nullOnly { + return nil, nullOnly, err + } + + schema.Delete("$defs") + schema.Delete("definitions") + if typeValue, ok := schema.Get("type"); ok && typeValue == "object" { + if _, hasProperties := schema.Get("properties"); !hasProperties { + schema.Set("properties", schemas.NewOrderedMap()) + } + } + + return schema, false, nil +} + +func sanitizeGigaChatSchemaComposition(schema *schemas.OrderedMap, root *schemas.OrderedMap, resolving map[string]bool, path string, keyword string) (*schemas.OrderedMap, bool, error) { + value, _ := schema.Get(keyword) + branches, ok := value.([]interface{}) + if !ok { + return nil, false, fmt.Errorf("%s.%s: expected an array", path, keyword) + } + + var selected *schemas.OrderedMap + for index, branch := range branches { + branchMap, ok := asGigaChatSchemaMap(branch) + if !ok { + return nil, false, fmt.Errorf("%s.%s[%d]: expected a JSON schema object", path, keyword, index) + } + branchCopy, err := cloneGigaChatSchemaMap(branchMap) + if err != nil { + return nil, false, fmt.Errorf("%s.%s[%d]: clone branch: %w", path, keyword, index, err) + } + sanitizedBranch, nullOnly, err := sanitizeGigaChatSchemaMap(branchCopy, root, resolving, fmt.Sprintf("%s.%s[%d]", path, keyword, index)) + if err != nil { + return nil, false, err + } + if nullOnly { + continue + } + if selected != nil { + return nil, false, fmt.Errorf("%s.%s: multiple non-null branches are not supported by GigaChat function schemas", path, keyword) + } + selected = sanitizedBranch + } + if selected == nil { + return nil, true, nil + } + + merged, err := cloneGigaChatSchemaMap(selected) + if err != nil { + return nil, false, fmt.Errorf("%s.%s: clone selected branch: %w", path, keyword, err) + } + var copyErr error + schema.Range(func(key string, value interface{}) bool { + if key == keyword || key == "$defs" || key == "definitions" { + return true + } + copied, err := cloneGigaChatSchemaValue(value) + if err != nil { + copyErr = err + return false + } + merged.Set(key, copied) + return true + }) + if copyErr != nil { + return nil, false, fmt.Errorf("%s.%s: merge branch siblings: %w", path, keyword, copyErr) + } + return sanitizeGigaChatSchemaMap(merged, root, resolving, path) +} + +func sanitizeGigaChatSchemaType(schema *schemas.OrderedMap, path string) (bool, error) { + value, ok := schema.Get("type") + if !ok { + return false, nil + } + + switch typed := value.(type) { + case string: + if typed == "null" { + return true, nil + } + return false, nil + case []interface{}: + nonNullTypes := make([]string, 0, len(typed)) + for _, item := range typed { + typeName, ok := item.(string) + if !ok { + return false, fmt.Errorf("%s.type: type arrays must contain strings", path) + } + if typeName != "null" { + nonNullTypes = append(nonNullTypes, typeName) + } + } + switch len(nonNullTypes) { + case 0: + return true, nil + case 1: + schema.Set("type", nonNullTypes[0]) + return false, nil + default: + return false, fmt.Errorf("%s.type: multiple non-null types are not supported by GigaChat function schemas", path) + } + default: + return false, fmt.Errorf("%s.type: expected a string or string array", path) + } +} + +func sanitizeGigaChatSchemaProperties(schema *schemas.OrderedMap, root *schemas.OrderedMap, resolving map[string]bool, path string) error { + value, ok := schema.Get("properties") + if !ok || value == nil { + return nil + } + + properties, ok := asGigaChatSchemaMap(value) + if !ok { + return fmt.Errorf("%s.properties: expected an object", path) + } + + sanitizedProperties := schemas.NewOrderedMap() + var sanitizeErr error + properties.Range(func(name string, propertyValue interface{}) bool { + propertySchema, ok := asGigaChatSchemaMap(propertyValue) + if !ok { + sanitizeErr = fmt.Errorf("%s.properties.%s: expected a JSON schema object", path, name) + return false + } + propertyCopy, err := cloneGigaChatSchemaMap(propertySchema) + if err != nil { + sanitizeErr = fmt.Errorf("%s.properties.%s: clone schema: %w", path, name, err) + return false + } + sanitizedProperty, nullOnly, err := sanitizeGigaChatSchemaMap(propertyCopy, root, resolving, path+".properties."+name) + if err != nil { + sanitizeErr = err + return false + } + if nullOnly { + sanitizeErr = fmt.Errorf("%s.properties.%s: null-only schemas are not supported by GigaChat function schemas", path, name) + return false + } + sanitizedProperties.Set(name, sanitizedProperty) + return true + }) + if sanitizeErr != nil { + return sanitizeErr + } + + schema.Set("properties", sanitizedProperties) + return nil +} + +func sanitizeGigaChatSchemaItems(schema *schemas.OrderedMap, root *schemas.OrderedMap, resolving map[string]bool, path string) error { + value, ok := schema.Get("items") + if !ok || value == nil { + return nil + } + + sanitized, nullOnly, err := sanitizeGigaChatNestedSchemaValue(value, root, resolving, path+".items") + if err != nil { + return err + } + if nullOnly { + return fmt.Errorf("%s.items: null-only schemas are not supported by GigaChat function schemas", path) + } + schema.Set("items", sanitized) + return nil +} + +func sanitizeGigaChatSchemaAdditionalProperties(schema *schemas.OrderedMap, root *schemas.OrderedMap, resolving map[string]bool, path string) error { + value, ok := schema.Get("additionalProperties") + if !ok || value == nil { + return nil + } + if _, ok := value.(bool); ok { + return nil + } + + sanitized, nullOnly, err := sanitizeGigaChatNestedSchemaValue(value, root, resolving, path+".additionalProperties") + if err != nil { + return err + } + if nullOnly { + return fmt.Errorf("%s.additionalProperties: null-only schemas are not supported by GigaChat function schemas", path) + } + schema.Set("additionalProperties", sanitized) + return nil +} + +func sanitizeGigaChatSchemaAllOf(schema *schemas.OrderedMap, root *schemas.OrderedMap, resolving map[string]bool, path string) (bool, error) { + value, ok := schema.Get("allOf") + if !ok || value == nil { + return false, nil + } + branches, ok := value.([]interface{}) + if !ok { + return false, fmt.Errorf("%s.allOf: expected an array", path) + } + + sanitizedBranches := make([]interface{}, 0, len(branches)) + for index, branch := range branches { + sanitized, nullOnly, err := sanitizeGigaChatNestedSchemaValue(branch, root, resolving, fmt.Sprintf("%s.allOf[%d]", path, index)) + if err != nil { + return false, err + } + if !nullOnly { + sanitizedBranches = append(sanitizedBranches, sanitized) + } + } + if len(branches) > 0 && len(sanitizedBranches) == 0 { + return true, nil + } + schema.Set("allOf", sanitizedBranches) + return false, nil +} + +func sanitizeGigaChatNestedSchemaValue(value interface{}, root *schemas.OrderedMap, resolving map[string]bool, path string) (interface{}, bool, error) { + if schemaMap, ok := asGigaChatSchemaMap(value); ok { + schemaCopy, err := cloneGigaChatSchemaMap(schemaMap) + if err != nil { + return nil, false, fmt.Errorf("%s: clone schema: %w", path, err) + } + return sanitizeGigaChatSchemaMap(schemaCopy, root, resolving, path) + } + + items, ok := value.([]interface{}) + if !ok { + return value, false, nil + } + + sanitizedItems := make([]interface{}, 0, len(items)) + for index, item := range items { + sanitized, nullOnly, err := sanitizeGigaChatNestedSchemaValue(item, root, resolving, fmt.Sprintf("%s[%d]", path, index)) + if err != nil { + return nil, false, err + } + if !nullOnly { + sanitizedItems = append(sanitizedItems, sanitized) + } + } + return sanitizedItems, len(sanitizedItems) == 0 && len(items) > 0, nil +} + +func lookupGigaChatLocalSchemaRef(root *schemas.OrderedMap, ref string) (*schemas.OrderedMap, bool) { + if root == nil || !strings.HasPrefix(ref, "#") { + return nil, false + } + if ref == "#" { + return root, true + } + if !strings.HasPrefix(ref, "#/") { + return nil, false + } + + var current interface{} = root + for _, rawToken := range strings.Split(strings.TrimPrefix(ref, "#/"), "/") { + token := decodeGigaChatJSONPointerToken(rawToken) + currentMap, ok := asGigaChatSchemaMap(current) + if !ok { + return nil, false + } + next, ok := currentMap.Get(token) + if !ok { + return nil, false + } + current = next + } + return asGigaChatSchemaMap(current) +} + +func decodeGigaChatJSONPointerToken(token string) string { + token = strings.ReplaceAll(token, "~1", "/") + return strings.ReplaceAll(token, "~0", "~") +} + +func asGigaChatSchemaMap(value interface{}) (*schemas.OrderedMap, bool) { + switch typed := value.(type) { + case *schemas.OrderedMap: + return typed, typed != nil + case schemas.OrderedMap: + return &typed, true + case map[string]interface{}: + return schemas.OrderedMapFromMap(typed), true + default: + return nil, false + } +} + +func cloneGigaChatSchemaMap(value *schemas.OrderedMap) (*schemas.OrderedMap, error) { + if value == nil { + return nil, nil + } + raw, err := schemas.MarshalSorted(value) + if err != nil { + return nil, err + } + var cloned schemas.OrderedMap + if err := schemas.Unmarshal(raw, &cloned); err != nil { + return nil, err + } + return &cloned, nil +} + +func cloneGigaChatSchemaValue(value interface{}) (interface{}, error) { + switch typed := value.(type) { + case *schemas.OrderedMap: + return cloneGigaChatSchemaMap(typed) + case schemas.OrderedMap: + return cloneGigaChatSchemaMap(&typed) + case map[string]interface{}: + copied := make(map[string]interface{}, len(typed)) + for key, item := range typed { + itemCopy, err := cloneGigaChatSchemaValue(item) + if err != nil { + return nil, err + } + copied[key] = itemCopy + } + return copied, nil + case []interface{}: + copied := make([]interface{}, len(typed)) + for index, item := range typed { + itemCopy, err := cloneGigaChatSchemaValue(item) + if err != nil { + return nil, err + } + copied[index] = itemCopy + } + return copied, nil + default: + return typed, nil + } +} diff --git a/core/providers/gigachat/tools.go b/core/providers/gigachat/tools.go new file mode 100644 index 00000000000..ce646a0ea52 --- /dev/null +++ b/core/providers/gigachat/tools.go @@ -0,0 +1,651 @@ +package gigachat + +import ( + "bytes" + "fmt" + "regexp" + "sort" + "strings" + + schemas "github.com/maximhq/bifrost/core/schemas" +) + +var gigaChatFunctionNamePattern = regexp.MustCompile(`^[A-Za-z_][A-Za-z0-9_]*$`) + +var gigaChatBuiltInFunctionNames = map[string]struct{}{ + "text2image": {}, + "get_file_content": {}, + "text2model3d": {}, +} + +const ( + gigaChatResponsesUserFunctionNamePrefix = "__bifrost_gigachat_user_" + gigaChatResponsesToolNameCodeInterpreter = "code_interpreter" + gigaChatResponsesToolNameImageGenerate = "image_generate" + gigaChatResponsesToolNameModel3DGenerate = "model_3d_generate" + gigaChatResponsesToolNameURLContentExtraction = "url_content_extraction" + gigaChatResponsesToolNameWebSearch = "web_search" + gigaChatResponsesToolTypeURLContentExtraction = "url_content_extraction" + gigaChatResponsesToolTypeModel3DGenerate = "model_3d_generate" + gigaChatResponsesSearchContextSizeFlagPrefix = "search_context_size:" + gigaChatResponsesUserLocationUserInfoField = "user_location" +) + +var gigaChatResponsesReservedFunctionNames = map[string]struct{}{ + "code_interpreter": {}, + "image_generate": {}, + "image_generation": {}, + "model_3d_generate": {}, + "url_content_extraction": {}, + "web_search": {}, + "web_search_preview": {}, +} + +// GigaChat built-ins are service-side functions with provider-specific side effects. +// They are intentionally rejected through neutral Bifrost tool fields until the provider has a scoped API for them. +func validateGigaChatFunctionName(name string) error { + trimmed := strings.TrimSpace(name) + if trimmed == "" { + return fmt.Errorf("function tool name is required") + } + if _, ok := gigaChatBuiltInFunctionNames[trimmed]; ok { + return fmt.Errorf("GigaChat built-in function %q is not supported through neutral function tools", trimmed) + } + if !gigaChatFunctionNamePattern.MatchString(trimmed) { + return fmt.Errorf("function tool name %q must start with a latin letter or underscore and contain only latin letters, digits, or underscores", trimmed) + } + return nil +} + +func validateGigaChatFunctionStrict(strict *bool) error { + if strict != nil && *strict { + return fmt.Errorf("function strict mode is not supported by GigaChat") + } + return nil +} + +func toGigaChatChatFunctions(tools []schemas.ChatTool) ([]GigaChatFunction, map[string]struct{}, error) { + if len(tools) == 0 { + return nil, nil, nil + } + + functions := make([]GigaChatFunction, 0, len(tools)) + functionNames := make(map[string]struct{}, len(tools)) + functionDefinitions := make(map[string]GigaChatFunction, len(tools)) + for index, tool := range tools { + if tool.Type != schemas.ChatToolTypeFunction { + return nil, nil, fmt.Errorf("tools[%d]: GigaChat chat completions support user-defined function tools only, got %q", index, tool.Type) + } + if tool.Function == nil { + return nil, nil, fmt.Errorf("tools[%d]: function tool definition is required", index) + } + name := strings.TrimSpace(tool.Function.Name) + if err := validateGigaChatFunctionName(name); err != nil { + return nil, nil, fmt.Errorf("tools[%d]: %w", index, err) + } + parameters, err := sanitizeGigaChatFunctionSchema(tool.Function.Parameters) + if err != nil { + return nil, nil, fmt.Errorf("tools[%d]: %w", index, err) + } + if err := validateGigaChatFunctionStrict(tool.Function.Strict); err != nil { + return nil, nil, fmt.Errorf("tools[%d]: %w", index, err) + } + if _, exists := functionNames[name]; exists { + sameDefinition, err := sameGigaChatToolDefinition(functionDefinitions[name], GigaChatFunction{ + Name: name, + Description: tool.Function.Description, + Parameters: parameters, + }) + if err != nil { + return nil, nil, fmt.Errorf("tools[%d]: compare duplicate function tool %q: %w", index, name, err) + } + if sameDefinition { + continue + } + return nil, nil, fmt.Errorf("tools[%d]: duplicate function tool name %q has a different definition", index, name) + } + + function := GigaChatFunction{ + Name: name, + Description: tool.Function.Description, + Parameters: parameters, + } + functionNames[name] = struct{}{} + functionDefinitions[name] = function + functions = append(functions, function) + } + + return functions, functionNames, nil +} + +type gigaChatResponsesToolsConversion struct { + Tools []GigaChatResponsesTool + UserInfo map[string]interface{} +} + +func toGigaChatChatFunctionCall(toolChoice *schemas.ChatToolChoice, functionNames map[string]struct{}) (interface{}, error) { + if toolChoice == nil { + return nil, nil + } + if toolChoice.ChatToolChoiceStr != nil { + switch strings.TrimSpace(*toolChoice.ChatToolChoiceStr) { + case "": + return nil, nil + case "auto": + if len(functionNames) == 0 { + return nil, fmt.Errorf("tool_choice auto requires at least one declared GigaChat function tool") + } + return "auto", nil + case "required", "any": + return forceSingleGigaChatChatFunctionChoice(functionNames, strings.TrimSpace(*toolChoice.ChatToolChoiceStr)) + case "none": + return "none", nil + default: + return nil, fmt.Errorf("tool_choice %q is not supported by GigaChat chat completions", *toolChoice.ChatToolChoiceStr) + } + } + if toolChoice.ChatToolChoiceStruct == nil { + return nil, nil + } + + choice := toolChoice.ChatToolChoiceStruct + switch choice.Type { + case schemas.ChatToolChoiceTypeFunction: + if choice.Function == nil || strings.TrimSpace(choice.Function.Name) == "" { + return nil, fmt.Errorf("tool_choice function name is required") + } + name := strings.TrimSpace(choice.Function.Name) + if _, ok := functionNames[name]; !ok { + return nil, fmt.Errorf("tool_choice function %q must match a declared GigaChat function tool", name) + } + return GigaChatFunctionCallChoice{Name: name}, nil + case schemas.ChatToolChoiceTypeAuto: + if len(functionNames) == 0 { + return nil, fmt.Errorf("tool_choice auto requires at least one declared GigaChat function tool") + } + return "auto", nil + case schemas.ChatToolChoiceTypeNone: + return "none", nil + case schemas.ChatToolChoiceTypeAny, schemas.ChatToolChoiceTypeRequired: + return forceSingleGigaChatChatFunctionChoice(functionNames, string(choice.Type)) + default: + return nil, fmt.Errorf("tool_choice type %q is not supported by GigaChat chat completions", choice.Type) + } +} + +func forceSingleGigaChatChatFunctionChoice(functionNames map[string]struct{}, choice string) (GigaChatFunctionCallChoice, error) { + if len(functionNames) == 0 { + return GigaChatFunctionCallChoice{}, fmt.Errorf("tool_choice %s requires at least one declared GigaChat function tool", choice) + } + if len(functionNames) > 1 { + return GigaChatFunctionCallChoice{}, fmt.Errorf("tool_choice %s cannot require an arbitrary GigaChat function when multiple function tools are declared", choice) + } + for name := range functionNames { + return GigaChatFunctionCallChoice{Name: name}, nil + } + return GigaChatFunctionCallChoice{}, fmt.Errorf("tool_choice %s requires at least one declared GigaChat function tool", choice) +} + +func toGigaChatResponsesTools(tools []schemas.ResponsesTool) (*gigaChatResponsesToolsConversion, error) { + converted := &gigaChatResponsesToolsConversion{} + if len(tools) == 0 { + return converted, nil + } + + specifications := make([]GigaChatResponsesFunctionSpecification, 0, len(tools)) + functionDefinitions := make(map[string]gigaChatResponsesFunctionDefinition, len(tools)) + functionsToolIndex := -1 + for index, tool := range tools { + switch { + case tool.Type == schemas.ResponsesToolTypeFunction: + specification, name, err := toGigaChatResponsesFunctionSpecification(index, tool) + if err != nil { + return nil, err + } + if existing, exists := functionDefinitions[specification.Name]; exists { + sameDefinition, err := sameGigaChatToolDefinition(existing.Specification, *specification) + if err != nil { + return nil, fmt.Errorf("tools[%d]: compare duplicate function tool %q: %w", index, name, err) + } + if sameDefinition { + continue + } + if existing.OriginalName == name { + return nil, fmt.Errorf("tools[%d]: duplicate function tool name %q has a different definition after GigaChat compatibility remapping", index, name) + } + return nil, fmt.Errorf("tools[%d]: duplicate function tool name %q conflicts with %q after GigaChat compatibility remapping", index, name, existing.OriginalName) + } + functionDefinitions[specification.Name] = gigaChatResponsesFunctionDefinition{ + OriginalName: name, + Specification: *specification, + } + if functionsToolIndex == -1 { + functionsToolIndex = len(converted.Tools) + converted.Tools = append(converted.Tools, GigaChatResponsesTool{}) + } + specifications = append(specifications, *specification) + case isGigaChatResponsesWebSearchToolType(tool.Type): + gigaChatTool, userInfo, err := toGigaChatResponsesWebSearchTool(index, tool) + if err != nil { + return nil, err + } + if len(userInfo) > 0 { + if converted.UserInfo != nil { + return nil, fmt.Errorf("tools[%d]: multiple web_search user_location configs are not supported by GigaChat Responses", index) + } + converted.UserInfo = userInfo + } + converted.Tools = append(converted.Tools, *gigaChatTool) + case tool.Type == schemas.ResponsesToolTypeCodeInterpreter: + gigaChatTool, err := toGigaChatResponsesCodeInterpreterTool(index, tool) + if err != nil { + return nil, err + } + converted.Tools = append(converted.Tools, *gigaChatTool) + case tool.Type == schemas.ResponsesToolTypeImageGeneration: + gigaChatTool, err := toGigaChatResponsesImageGenerateTool(index, tool) + if err != nil { + return nil, err + } + converted.Tools = append(converted.Tools, *gigaChatTool) + case tool.Type == schemas.ResponsesToolTypeWebFetch || string(tool.Type) == gigaChatResponsesToolTypeURLContentExtraction: + gigaChatTool, err := toGigaChatResponsesURLContentExtractionTool(index, tool) + if err != nil { + return nil, err + } + converted.Tools = append(converted.Tools, *gigaChatTool) + case string(tool.Type) == gigaChatResponsesToolTypeModel3DGenerate: + gigaChatTool, err := toGigaChatResponsesModel3DGenerateTool(index, tool) + if err != nil { + return nil, err + } + converted.Tools = append(converted.Tools, *gigaChatTool) + default: + return nil, fmt.Errorf("tools[%d]: GigaChat Responses does not support tool type %q", index, tool.Type) + } + } + + if len(specifications) > 0 { + converted.Tools[functionsToolIndex].Functions = &GigaChatResponsesFunctionsTool{ + Specifications: specifications, + } + } + + return converted, nil +} + +type gigaChatResponsesFunctionDefinition struct { + OriginalName string + Specification GigaChatResponsesFunctionSpecification +} + +func toGigaChatResponsesFunctionSpecification(index int, tool schemas.ResponsesTool) (*GigaChatResponsesFunctionSpecification, string, error) { + if tool.Name == nil || strings.TrimSpace(*tool.Name) == "" { + return nil, "", fmt.Errorf("tools[%d]: function tool name is required", index) + } + name := strings.TrimSpace(*tool.Name) + if strings.HasPrefix(name, gigaChatResponsesUserFunctionNamePrefix) { + return nil, "", fmt.Errorf("tools[%d]: function tool name %q uses a GigaChat compatibility-reserved prefix", index, name) + } + if err := validateGigaChatFunctionName(name); err != nil { + return nil, "", fmt.Errorf("tools[%d]: %w", index, err) + } + if tool.ResponsesToolFunction == nil { + return nil, "", fmt.Errorf("tools[%d]: function tool definition is required", index) + } + parameters, err := sanitizeGigaChatFunctionSchema(tool.ResponsesToolFunction.Parameters) + if err != nil { + return nil, "", fmt.Errorf("tools[%d]: %w", index, err) + } + if err := validateGigaChatFunctionStrict(tool.ResponsesToolFunction.Strict); err != nil { + return nil, "", fmt.Errorf("tools[%d]: %w", index, err) + } + + gigaChatName := toGigaChatResponsesFunctionName(name) + + return &GigaChatResponsesFunctionSpecification{ + Name: gigaChatName, + Description: tool.Description, + Parameters: parameters, + }, name, nil +} + +func sameGigaChatToolDefinition(left interface{}, right interface{}) (bool, error) { + leftRaw, err := schemas.MarshalSorted(left) + if err != nil { + return false, err + } + rightRaw, err := schemas.MarshalSorted(right) + if err != nil { + return false, err + } + return bytes.Equal(leftRaw, rightRaw), nil +} + +func toGigaChatResponsesFunctionName(name string) string { + trimmed := strings.TrimSpace(name) + if _, ok := gigaChatResponsesReservedFunctionNames[trimmed]; ok { + return gigaChatResponsesUserFunctionNamePrefix + trimmed + } + return trimmed +} + +func toBifrostGigaChatResponsesFunctionName(name string) string { + trimmed := strings.TrimSpace(name) + if !strings.HasPrefix(trimmed, gigaChatResponsesUserFunctionNamePrefix) { + return trimmed + } + original := strings.TrimPrefix(trimmed, gigaChatResponsesUserFunctionNamePrefix) + if _, ok := gigaChatResponsesReservedFunctionNames[original]; ok { + return original + } + return trimmed +} + +func isGigaChatResponsesWebSearchToolType(toolType schemas.ResponsesToolType) bool { + value := strings.TrimSpace(string(toolType)) + return value == string(schemas.ResponsesToolTypeWebSearch) || + value == string(schemas.ResponsesToolTypeWebSearchPreview) || + strings.HasPrefix(value, "web_search_") +} + +func toGigaChatResponsesWebSearchTool(index int, tool schemas.ResponsesTool) (*GigaChatResponsesTool, map[string]interface{}, error) { + webSearch := &GigaChatResponsesWebSearchTool{} + var userInfo map[string]interface{} + + if tool.ResponsesToolWebSearch != nil { + if tool.ResponsesToolWebSearch.Filters != nil { + return nil, nil, fmt.Errorf("tools[%d]: web_search filters are not supported by GigaChat Responses", index) + } + if len(tool.ResponsesToolWebSearch.SearchContentTypes) > 0 { + return nil, nil, fmt.Errorf("tools[%d]: web_search search_content_types are not supported by GigaChat Responses", index) + } + if tool.ResponsesToolWebSearch.ExternalWebAccess != nil { + return nil, nil, fmt.Errorf("tools[%d]: web_search external_web_access is not supported by GigaChat Responses", index) + } + if tool.ResponsesToolWebSearch.MaxUses != nil { + return nil, nil, fmt.Errorf("tools[%d]: web_search max_uses is not supported by GigaChat Responses", index) + } + if tool.ResponsesToolWebSearch.SearchContextSize != nil && strings.TrimSpace(*tool.ResponsesToolWebSearch.SearchContextSize) != "" { + webSearch.Flags = append(webSearch.Flags, gigaChatResponsesSearchContextSizeFlagPrefix+strings.TrimSpace(*tool.ResponsesToolWebSearch.SearchContextSize)) + } + if tool.ResponsesToolWebSearch.UserLocation != nil { + userInfo = toGigaChatResponsesUserInfo(tool.ResponsesToolWebSearch.UserLocation) + } + } + if tool.ResponsesToolWebSearchPreview != nil { + if tool.ResponsesToolWebSearchPreview.SearchContextSize != nil && strings.TrimSpace(*tool.ResponsesToolWebSearchPreview.SearchContextSize) != "" { + webSearch.Flags = append(webSearch.Flags, gigaChatResponsesSearchContextSizeFlagPrefix+strings.TrimSpace(*tool.ResponsesToolWebSearchPreview.SearchContextSize)) + } + if tool.ResponsesToolWebSearchPreview.UserLocation != nil { + userInfo = toGigaChatResponsesUserInfo(tool.ResponsesToolWebSearchPreview.UserLocation) + } + } + + return &GigaChatResponsesTool{WebSearch: webSearch}, userInfo, nil +} + +func toGigaChatResponsesUserInfo(location *schemas.ResponsesToolWebSearchUserLocation) map[string]interface{} { + if location == nil { + return nil + } + userLocation := make(map[string]interface{}) + if location.Type != nil && strings.TrimSpace(*location.Type) != "" { + userLocation["type"] = strings.TrimSpace(*location.Type) + } + if location.City != nil && strings.TrimSpace(*location.City) != "" { + userLocation["city"] = strings.TrimSpace(*location.City) + } + if location.Country != nil && strings.TrimSpace(*location.Country) != "" { + userLocation["country"] = strings.TrimSpace(*location.Country) + } + if location.Region != nil && strings.TrimSpace(*location.Region) != "" { + userLocation["region"] = strings.TrimSpace(*location.Region) + } + if location.Timezone != nil && strings.TrimSpace(*location.Timezone) != "" { + userLocation["timezone"] = strings.TrimSpace(*location.Timezone) + } + if len(userLocation) == 0 { + return nil + } + return map[string]interface{}{gigaChatResponsesUserLocationUserInfoField: userLocation} +} + +func toGigaChatResponsesCodeInterpreterTool(index int, tool schemas.ResponsesTool) (*GigaChatResponsesTool, error) { + config, err := toGigaChatResponsesToolConfigMap(tool.ResponsesToolCodeInterpreter) + if err != nil { + return nil, fmt.Errorf("tools[%d]: code_interpreter config is invalid: %w", index, err) + } + return &GigaChatResponsesTool{CodeInterpreter: config}, nil +} + +func toGigaChatResponsesImageGenerateTool(index int, tool schemas.ResponsesTool) (*GigaChatResponsesTool, error) { + config, err := toGigaChatResponsesToolConfigMap(tool.ResponsesToolImageGeneration) + if err != nil { + return nil, fmt.Errorf("tools[%d]: image_generation config is invalid: %w", index, err) + } + return &GigaChatResponsesTool{ImageGenerate: config}, nil +} + +func toGigaChatResponsesURLContentExtractionTool(index int, tool schemas.ResponsesTool) (*GigaChatResponsesTool, error) { + config, err := toGigaChatResponsesToolConfigMap(tool.ResponsesToolWebFetch) + if err != nil { + return nil, fmt.Errorf("tools[%d]: url_content_extraction config is invalid: %w", index, err) + } + addGigaChatResponsesCommonToolFields(config, tool) + return &GigaChatResponsesTool{URLContentExtraction: config}, nil +} + +func toGigaChatResponsesModel3DGenerateTool(index int, tool schemas.ResponsesTool) (*GigaChatResponsesTool, error) { + config, err := toGigaChatResponsesToolConfigMap(nil) + if err != nil { + return nil, fmt.Errorf("tools[%d]: model_3d_generate config is invalid: %w", index, err) + } + addGigaChatResponsesCommonToolFields(config, tool) + return &GigaChatResponsesTool{Model3DGenerate: config}, nil +} + +func toGigaChatResponsesToolConfigMap(value interface{}) (map[string]interface{}, error) { + config := map[string]interface{}{} + if value == nil { + return config, nil + } + + raw, err := schemas.MarshalSorted(value) + if err != nil { + return nil, err + } + var fields map[string]interface{} + if err := schemas.Unmarshal(raw, &fields); err != nil { + return nil, err + } + for name, fieldValue := range fields { + if fieldValue != nil { + config[name] = fieldValue + } + } + return config, nil +} + +func addGigaChatResponsesCommonToolFields(config map[string]interface{}, tool schemas.ResponsesTool) { + if config == nil { + return + } + if tool.Name != nil && strings.TrimSpace(*tool.Name) != "" { + config["name"] = strings.TrimSpace(*tool.Name) + } + if tool.Description != nil && strings.TrimSpace(*tool.Description) != "" { + config["description"] = strings.TrimSpace(*tool.Description) + } +} + +func toGigaChatResponsesToolConfig(toolChoice *schemas.ResponsesToolChoice, tools []schemas.ResponsesTool) (*GigaChatResponsesToolConfig, error) { + if toolChoice == nil { + return nil, nil + } + targets := newGigaChatResponsesToolChoiceTargets(tools) + if toolChoice.ResponsesToolChoiceStr != nil { + switch strings.TrimSpace(*toolChoice.ResponsesToolChoiceStr) { + case "": + return nil, nil + case "auto": + if !targets.HasTools() { + return nil, fmt.Errorf("tool_choice auto requires at least one declared GigaChat tool") + } + return &GigaChatResponsesToolConfig{Mode: "auto"}, nil + case "none": + return &GigaChatResponsesToolConfig{Mode: "none"}, nil + case "required", "any": + return forceSingleGigaChatResponsesToolConfig(targets, strings.TrimSpace(*toolChoice.ResponsesToolChoiceStr)) + default: + return nil, fmt.Errorf("tool_choice %q is not supported by GigaChat Responses", *toolChoice.ResponsesToolChoiceStr) + } + } + if toolChoice.ResponsesToolChoiceStruct == nil { + return nil, nil + } + + choice := toolChoice.ResponsesToolChoiceStruct + switch choice.Type { + case schemas.ResponsesToolChoiceTypeFunction: + if choice.Name == nil || strings.TrimSpace(*choice.Name) == "" { + return nil, fmt.Errorf("tool_choice function name is required") + } + name := strings.TrimSpace(*choice.Name) + gigaChatName, ok := targets.Functions[name] + if !ok { + return nil, fmt.Errorf("tool_choice function %q must match a declared GigaChat function tool", name) + } + return &GigaChatResponsesToolConfig{ + Mode: "forced", + FunctionName: &gigaChatName, + }, nil + case schemas.ResponsesToolChoiceTypeAuto: + if !targets.HasTools() { + return nil, fmt.Errorf("tool_choice auto requires at least one declared GigaChat tool") + } + return &GigaChatResponsesToolConfig{Mode: "auto"}, nil + case schemas.ResponsesToolChoiceTypeNone: + return &GigaChatResponsesToolConfig{Mode: "none"}, nil + case schemas.ResponsesToolChoiceTypeAny, schemas.ResponsesToolChoiceTypeRequired: + return forceSingleGigaChatResponsesToolConfig(targets, string(choice.Type)) + case schemas.ResponsesToolChoiceTypeAllowedTools: + return nil, fmt.Errorf("tool_choice type %q is not supported by GigaChat Responses because tool_config supports one forced tool_name or function_name, not an allowed tools set", choice.Type) + case schemas.ResponsesToolChoiceTypeFileSearch, schemas.ResponsesToolChoiceTypeComputerUsePreview, schemas.ResponsesToolChoiceTypeMCP, schemas.ResponsesToolChoiceTypeCustom: + return nil, fmt.Errorf("tool_choice type %q is not supported by GigaChat Responses", choice.Type) + default: + toolName, ok := targets.BuiltIns[gigaChatResponsesToolChoiceTypeToBuiltInName(choice.Type)] + if !ok { + return nil, fmt.Errorf("tool_choice type %q must match a declared GigaChat built-in tool", choice.Type) + } + return &GigaChatResponsesToolConfig{ + Mode: "forced", + ToolName: &toolName, + }, nil + } +} + +func forceSingleGigaChatResponsesToolConfig(targets gigaChatResponsesToolChoiceTargets, choice string) (*GigaChatResponsesToolConfig, error) { + targetCount := len(targets.Functions) + len(targets.BuiltIns) + if targetCount == 0 { + return nil, fmt.Errorf("tool_choice %s requires at least one declared GigaChat tool", choice) + } + if targetCount > 1 { + return nil, fmt.Errorf("tool_choice %s cannot require an arbitrary GigaChat tool when multiple tools are declared", choice) + } + for _, name := range targets.Functions { + return &GigaChatResponsesToolConfig{ + Mode: "forced", + FunctionName: &name, + }, nil + } + for _, name := range targets.BuiltIns { + return &GigaChatResponsesToolConfig{ + Mode: "forced", + ToolName: &name, + }, nil + } + return nil, fmt.Errorf("tool_choice %s requires at least one declared GigaChat tool", choice) +} + +type gigaChatResponsesToolChoiceTargets struct { + Functions map[string]string + BuiltIns map[string]string +} + +func newGigaChatResponsesToolChoiceTargets(tools []schemas.ResponsesTool) gigaChatResponsesToolChoiceTargets { + targets := gigaChatResponsesToolChoiceTargets{ + Functions: make(map[string]string, len(tools)), + BuiltIns: make(map[string]string, len(tools)), + } + + for _, tool := range tools { + if tool.Type == schemas.ResponsesToolTypeFunction && tool.Name != nil { + name := strings.TrimSpace(*tool.Name) + if name != "" { + targets.Functions[name] = toGigaChatResponsesFunctionName(name) + } + continue + } + if toolName, ok := gigaChatResponsesToolTypeToBuiltInName(tool.Type); ok { + targets.BuiltIns[toolName] = toolName + } + } + return targets +} + +func (targets gigaChatResponsesToolChoiceTargets) HasTools() bool { + return len(targets.Functions) > 0 || len(targets.BuiltIns) > 0 +} + +func gigaChatResponsesToolTypeToBuiltInName(toolType schemas.ResponsesToolType) (string, bool) { + switch { + case isGigaChatResponsesWebSearchToolType(toolType): + return gigaChatResponsesToolNameWebSearch, true + case toolType == schemas.ResponsesToolTypeCodeInterpreter: + return gigaChatResponsesToolNameCodeInterpreter, true + case toolType == schemas.ResponsesToolTypeImageGeneration: + return gigaChatResponsesToolNameImageGenerate, true + case toolType == schemas.ResponsesToolTypeWebFetch || string(toolType) == gigaChatResponsesToolTypeURLContentExtraction: + return gigaChatResponsesToolNameURLContentExtraction, true + case string(toolType) == gigaChatResponsesToolTypeModel3DGenerate: + return gigaChatResponsesToolNameModel3DGenerate, true + default: + return "", false + } +} + +func gigaChatResponsesToolChoiceTypeToBuiltInName(choiceType schemas.ResponsesToolChoiceType) string { + value := strings.TrimSpace(string(choiceType)) + switch { + case value == string(schemas.ResponsesToolChoiceTypeCodeInterpreter): + return gigaChatResponsesToolNameCodeInterpreter + case value == string(schemas.ResponsesToolChoiceTypeImageGeneration): + return gigaChatResponsesToolNameImageGenerate + case value == string(schemas.ResponsesToolChoiceTypeWebSearchPreview) || + value == string(schemas.ResponsesToolTypeWebSearch) || + strings.HasPrefix(value, string(schemas.ResponsesToolTypeWebSearch)+"_"): + return gigaChatResponsesToolNameWebSearch + case value == string(schemas.ResponsesToolTypeWebFetch) || value == gigaChatResponsesToolTypeURLContentExtraction: + return gigaChatResponsesToolNameURLContentExtraction + case value == gigaChatResponsesToolTypeModel3DGenerate: + return gigaChatResponsesToolNameModel3DGenerate + default: + return "" + } +} + +func unsupportedGigaChatToolControlExtraParams(extraParams map[string]interface{}, names ...string) []string { + if len(extraParams) == 0 { + return nil + } + + unsupported := make([]string, 0, len(names)) + for _, name := range names { + if _, ok := extraParams[name]; ok { + unsupported = append(unsupported, "extra_params."+name) + } + } + sort.Strings(unsupported) + return unsupported +} diff --git a/core/providers/gigachat/types.go b/core/providers/gigachat/types.go new file mode 100644 index 00000000000..38a9d752d7b --- /dev/null +++ b/core/providers/gigachat/types.go @@ -0,0 +1,627 @@ +// Package gigachat implements the GigaChat LLM provider. +package gigachat + +import ( + "bytes" + "encoding/json" + + schemas "github.com/maximhq/bifrost/core/schemas" +) + +// # AUTH TYPES + +// GigaChatTokenResponse is returned by the GigaChat OAuth endpoint. +type GigaChatTokenResponse struct { + AccessToken string `json:"access_token"` + ExpiresAt int64 `json:"expires_at"` +} + +// GigaChatPasswordTokenResponse is returned by the SDK-backed password auth endpoint. +type GigaChatPasswordTokenResponse struct { + Token string `json:"tok"` + ExpiresAt int64 `json:"exp"` +} + +// # CHAT TYPES + +// GigaChatChatRequest is the v1 chat completions request body. +type GigaChatChatRequest struct { + Model string `json:"model"` + Messages []GigaChatChatMessage `json:"messages"` + Temperature *float64 `json:"temperature,omitempty"` + TopP *float64 `json:"top_p,omitempty"` + MaxTokens *int `json:"max_tokens,omitempty"` + N *int `json:"n,omitempty"` + Stop []string `json:"stop,omitempty"` + Stream *bool `json:"stream,omitempty"` + ReasoningEffort *string `json:"reasoning_effort,omitempty"` + ResponseFormat interface{} `json:"response_format,omitempty"` + FunctionCall interface{} `json:"function_call,omitempty"` + Functions []GigaChatFunction `json:"functions,omitempty"` + ExtraParams map[string]interface{} `json:"-"` +} + +// GetExtraParams returns provider-specific passthrough fields. +func (request *GigaChatChatRequest) GetExtraParams() map[string]interface{} { + if request == nil || request.ExtraParams == nil { + return make(map[string]interface{}, 0) + } + return request.ExtraParams +} + +// GigaChatChatMessage is a GigaChat v1 chat message. +type GigaChatChatMessage struct { + Role string `json:"role,omitempty"` + Content *schemas.ChatMessageContent `json:"content,omitempty"` + Attachments []string `json:"attachments,omitempty"` + Name *string `json:"name,omitempty"` + Reasoning *string `json:"reasoning_content,omitempty"` + FunctionCall *GigaChatFunctionCall `json:"function_call,omitempty"` + FunctionsStateID *string `json:"functions_state_id,omitempty"` +} + +// UnmarshalJSON accepts both GigaChat's reasoning_content field and legacy +// reasoning-shaped payloads while preserving reasoning_content for outbound JSON. +func (message *GigaChatChatMessage) UnmarshalJSON(data []byte) error { + type Alias GigaChatChatMessage + var aux struct { + Alias + LegacyReasoning *string `json:"reasoning,omitempty"` + } + if err := json.Unmarshal(data, &aux); err != nil { + return err + } + + *message = GigaChatChatMessage(aux.Alias) + if message.Reasoning == nil && aux.LegacyReasoning != nil { + message.Reasoning = aux.LegacyReasoning + } + return nil +} + +// GigaChatFunctionCall is the legacy GigaChat function-call shape. +type GigaChatFunctionCall struct { + Name string `json:"name,omitempty"` + Arguments json.RawMessage `json:"arguments,omitempty"` +} + +// GigaChatFunctionCallChoice forces a specific GigaChat function call. +type GigaChatFunctionCallChoice struct { + Name string `json:"name"` +} + +// GigaChatFunction describes a client-defined function for GigaChat function calling. +type GigaChatFunction struct { + Name string `json:"name"` + Description *string `json:"description,omitempty"` + Parameters *schemas.ToolFunctionParameters `json:"parameters,omitempty"` + FewShotExamples []map[string]interface{} `json:"few_shot_examples,omitempty"` + ReturnParameters map[string]interface{} `json:"return_parameters,omitempty"` +} + +// GigaChatChatResponse is the v1 chat completions response body. +type GigaChatChatResponse struct { + ID string `json:"id,omitempty"` + Choices []GigaChatChatChoice `json:"choices,omitempty"` + Created int `json:"created,omitempty"` + Model string `json:"model,omitempty"` + Object string `json:"object,omitempty"` + SystemFingerprint string `json:"system_fingerprint,omitempty"` + Usage *GigaChatChatUsage `json:"usage,omitempty"` + ExtraParams map[string]interface{} `json:"-"` +} + +// GigaChatChatChoice is a single v1 chat completion choice. +type GigaChatChatChoice struct { + Index int `json:"index"` + Message *GigaChatChatMessage `json:"message,omitempty"` + FinishReason *string `json:"finish_reason,omitempty"` + LogProbs *schemas.BifrostLogProbs `json:"logprobs,omitempty"` +} + +// GigaChatChatStreamResponse is a v1 chat completions SSE chunk. +type GigaChatChatStreamResponse struct { + ID string `json:"id,omitempty"` + Choices []GigaChatChatStreamChoice `json:"choices,omitempty"` + Created int `json:"created,omitempty"` + Model string `json:"model,omitempty"` + Object string `json:"object,omitempty"` + SystemFingerprint string `json:"system_fingerprint,omitempty"` + Usage *GigaChatChatUsage `json:"usage,omitempty"` + ExtraParams map[string]interface{} `json:"-"` +} + +// GigaChatChatStreamChoice is a single streaming choice. +type GigaChatChatStreamChoice struct { + Index int `json:"index"` + Delta *GigaChatChatStreamDelta `json:"delta,omitempty"` + FinishReason *string `json:"finish_reason,omitempty"` + LogProbs *schemas.BifrostLogProbs `json:"logprobs,omitempty"` +} + +// GigaChatChatStreamDelta is the partial assistant message in an SSE chunk. +type GigaChatChatStreamDelta struct { + Role *string `json:"role,omitempty"` + Content *string `json:"content,omitempty"` + Reasoning *string `json:"reasoning_content,omitempty"` + FunctionCall *GigaChatFunctionCall `json:"function_call,omitempty"` + FunctionsStateID *string `json:"functions_state_id,omitempty"` +} + +// UnmarshalJSON accepts both reasoning_content and reasoning stream fields. +func (delta *GigaChatChatStreamDelta) UnmarshalJSON(data []byte) error { + type Alias GigaChatChatStreamDelta + var aux struct { + Alias + LegacyReasoning *string `json:"reasoning,omitempty"` + } + if err := json.Unmarshal(data, &aux); err != nil { + return err + } + + *delta = GigaChatChatStreamDelta(aux.Alias) + if delta.Reasoning == nil && aux.LegacyReasoning != nil { + delta.Reasoning = aux.LegacyReasoning + } + return nil +} + +// GigaChatChatUsage is token usage returned by GigaChat chat completions. +type GigaChatChatUsage struct { + PromptTokens int `json:"prompt_tokens,omitempty"` + CompletionTokens int `json:"completion_tokens,omitempty"` + TotalTokens int `json:"total_tokens,omitempty"` + PrecachedPromptTokens int `json:"precached_prompt_tokens,omitempty"` + InputTokens int `json:"input_tokens,omitempty"` + OutputTokens int `json:"output_tokens,omitempty"` + InputTokensDetails *GigaChatTokenDetails `json:"input_tokens_details,omitempty"` +} + +// GigaChatTokenDetails is the token breakdown shape used by GigaChat v2. +type GigaChatTokenDetails struct { + CachedTokens int `json:"cached_tokens,omitempty"` + CachedReadTokens int `json:"cached_read_tokens,omitempty"` +} + +// # MODELS TYPES + +// GigaChatListModelsResponse is the v1 models list response body. +type GigaChatListModelsResponse struct { + Object string `json:"object"` + Data []GigaChatModel `json:"data"` +} + +// GigaChatModel is a single model descriptor returned by GigaChat. +type GigaChatModel struct { + ID string `json:"id"` + Object string `json:"object,omitempty"` + OwnedBy string `json:"owned_by,omitempty"` + Type string `json:"type,omitempty"` +} + +// # FILES TYPES + +// GigaChatUploadedFile is a file metadata object returned by GigaChat. +type GigaChatUploadedFile struct { + ID string `json:"id"` + Object string `json:"object,omitempty"` + Bytes int64 `json:"bytes"` + CreatedAt int64 `json:"created_at"` + Filename string `json:"filename"` + Purpose string `json:"purpose"` + AccessPolicy *string `json:"access_policy,omitempty"` +} + +// GigaChatUploadedFiles is a list wrapper for GigaChat file metadata. +type GigaChatUploadedFiles struct { + Data []GigaChatUploadedFile `json:"data"` +} + +// GigaChatDeletedFile is returned by GigaChat after deleting a file. +type GigaChatDeletedFile struct { + ID string `json:"id"` + Deleted bool `json:"deleted"` +} + +// GigaChatFileContent contains base64-encoded file content. +type GigaChatFileContent struct { + Content string `json:"content"` +} + +// # BATCH TYPES + +// GigaChatBatchMethod selects the target operation for GigaChat batch execution. +type GigaChatBatchMethod string + +const ( + GigaChatBatchMethodChatCompletions GigaChatBatchMethod = "chat_completions" + GigaChatBatchMethodResponses GigaChatBatchMethod = "responses" + GigaChatBatchMethodEmbedder GigaChatBatchMethod = "embedder" +) + +// GigaChatBatchStatus is the lifecycle state returned by GigaChat batch APIs. +type GigaChatBatchStatus string + +const ( + GigaChatBatchStatusCreated GigaChatBatchStatus = "created" + GigaChatBatchStatusInProgress GigaChatBatchStatus = "in_progress" + GigaChatBatchStatusCompleted GigaChatBatchStatus = "completed" +) + +// GigaChatBatchRequestCounts tracks processed rows in a GigaChat batch job. +type GigaChatBatchRequestCounts struct { + Total int `json:"total,omitempty"` + Completed int `json:"completed,omitempty"` + Failed int `json:"failed,omitempty"` +} + +// GigaChatBatch is a batch metadata object returned by GigaChat. +type GigaChatBatch struct { + ID string `json:"id"` + Object string `json:"object,omitempty"` + Method GigaChatBatchMethod `json:"method,omitempty"` + Status GigaChatBatchStatus `json:"status,omitempty"` + RequestCounts *GigaChatBatchRequestCounts `json:"request_counts,omitempty"` + InputFileID *string `json:"input_file_id,omitempty"` + OutputFileID *string `json:"output_file_id,omitempty"` + ResultFileID *string `json:"result_file_id,omitempty"` + ErrorFileID *string `json:"error_file_id,omitempty"` + CompletionWindow string `json:"completion_window,omitempty"` + CreatedAt int64 `json:"created_at,omitempty"` + UpdatedAt *int64 `json:"updated_at,omitempty"` + CompletedAt *int64 `json:"completed_at,omitempty"` +} + +// GigaChatBatches is a list wrapper for GigaChat batch metadata. +type GigaChatBatches struct { + Data []GigaChatBatch `json:"data"` +} + +// UnmarshalJSON accepts both the wrapper shape returned by some GigaChat +// environments and the root list shape documented for empty task lists. +func (batches *GigaChatBatches) UnmarshalJSON(data []byte) error { + trimmed := bytes.TrimSpace(data) + if len(trimmed) > 0 && trimmed[0] == '[' { + var items []GigaChatBatch + if err := json.Unmarshal(trimmed, &items); err != nil { + return err + } + batches.Data = items + return nil + } + + type Alias GigaChatBatches + var aux Alias + if err := json.Unmarshal(trimmed, &aux); err != nil { + return err + } + *batches = GigaChatBatches(aux) + return nil +} + +// GigaChatBatchInputRow is a single JSONL row accepted by GigaChat batches. +type GigaChatBatchInputRow struct { + ID string `json:"id"` + Request json.RawMessage `json:"request"` +} + +// GigaChatBatchResultRow is a single JSONL row returned by GigaChat batches. +type GigaChatBatchResultRow struct { + ID string `json:"id,omitempty"` + CustomID string `json:"custom_id,omitempty"` + Response *schemas.BatchResultResponse `json:"response,omitempty"` + Result *schemas.BatchResultData `json:"result,omitempty"` + Error *schemas.BatchResultError `json:"error,omitempty"` +} + +// # EMBEDDING TYPES + +// GigaChatEmbeddingRequest is the v1 embeddings request body. +type GigaChatEmbeddingRequest struct { + Model string `json:"model"` + Input *schemas.EmbeddingInput `json:"input"` + + ExtraParams map[string]interface{} `json:"-"` +} + +// GetExtraParams returns provider-specific passthrough fields. +func (request *GigaChatEmbeddingRequest) GetExtraParams() map[string]interface{} { + if request == nil || request.ExtraParams == nil { + return make(map[string]interface{}, 0) + } + return request.ExtraParams +} + +// GigaChatEmbeddingResponse is the v1 embeddings response body. +type GigaChatEmbeddingResponse struct { + Object string `json:"object"` + Data []GigaChatEmbeddingData `json:"data"` + Model string `json:"model"` + Usage *GigaChatEmbeddingUsage `json:"usage,omitempty"` +} + +// GigaChatEmbeddingData is a single embedding vector returned by GigaChat. +type GigaChatEmbeddingData struct { + Object string `json:"object,omitempty"` + Embedding []float64 `json:"embedding"` + Index int `json:"index"` + Usage *GigaChatEmbeddingUsage `json:"usage,omitempty"` +} + +// GigaChatEmbeddingUsage describes token usage for embedding generation. +type GigaChatEmbeddingUsage struct { + PromptTokens int `json:"prompt_tokens,omitempty"` + TotalTokens int `json:"total_tokens,omitempty"` +} + +// # COUNT TOKENS TYPES + +// GigaChatCountTokensRequest is the v1 /tokens/count request body. +type GigaChatCountTokensRequest struct { + Model string `json:"model"` + Input []string `json:"input"` + ExtraParams map[string]interface{} `json:"-"` +} + +// GetExtraParams returns provider-specific passthrough fields. +func (request *GigaChatCountTokensRequest) GetExtraParams() map[string]interface{} { + if request == nil || request.ExtraParams == nil { + return make(map[string]interface{}, 0) + } + return request.ExtraParams +} + +// GigaChatCountTokensResponse accepts the documented SDK/API response shapes for +// /tokens/count: a root array, a {data:[...]} wrapper, or an aggregate object. +type GigaChatCountTokensResponse struct { + Object string `json:"object,omitempty"` + Model string `json:"model,omitempty"` + Data []GigaChatCountTokensItem `json:"data,omitempty"` + Tokens *int `json:"tokens,omitempty"` + Characters *int `json:"characters,omitempty"` + Items []GigaChatCountTokensItem `json:"-"` +} + +func (response *GigaChatCountTokensResponse) UnmarshalJSON(data []byte) error { + var items []GigaChatCountTokensItem + if err := json.Unmarshal(data, &items); err == nil { + response.Items = items + response.Data = items + return nil + } + + type Alias GigaChatCountTokensResponse + var object Alias + if err := json.Unmarshal(data, &object); err != nil { + return err + } + + *response = GigaChatCountTokensResponse(object) + if len(response.Data) > 0 { + response.Items = response.Data + return nil + } + if response.Tokens != nil { + item := GigaChatCountTokensItem{Tokens: *response.Tokens} + if response.Characters != nil { + item.Characters = *response.Characters + } + response.Items = []GigaChatCountTokensItem{item} + } + return nil +} + +// GigaChatCountTokensItem is one token count result for one input string. +type GigaChatCountTokensItem struct { + Tokens int `json:"tokens"` + Characters int `json:"characters,omitempty"` +} + +// # RESPONSES TYPES + +// GigaChatResponsesRequest is the v2 chat completions request body used for Bifrost Responses. +type GigaChatResponsesRequest struct { + Model string `json:"model,omitempty"` + Messages []GigaChatResponsesMessage `json:"messages"` + AssistantID *string `json:"assistant_id,omitempty"` + ToolsStateID *string `json:"tools_state_id,omitempty"` + ModelOptions *GigaChatResponsesModelOptions `json:"model_options,omitempty"` + FilterConfig map[string]interface{} `json:"filter_config,omitempty"` + Storage interface{} `json:"storage,omitempty"` + RankerOptions map[string]interface{} `json:"ranker_options,omitempty"` + ToolConfig *GigaChatResponsesToolConfig `json:"tool_config,omitempty"` + Tools []GigaChatResponsesTool `json:"tools,omitempty"` + UserInfo map[string]interface{} `json:"user_info,omitempty"` + Stream *bool `json:"stream,omitempty"` + DisableFilter *bool `json:"disable_filter,omitempty"` + Flags []string `json:"flags,omitempty"` + ExtraParams map[string]interface{} `json:"-"` +} + +// GigaChatResponsesStorage configures v2 thread storage. +type GigaChatResponsesStorage struct { + Limit *int `json:"limit,omitempty"` + ThreadID *string `json:"thread_id,omitempty"` + Metadata map[string]interface{} `json:"metadata,omitempty"` +} + +// GetExtraParams returns provider-specific passthrough fields. +func (request *GigaChatResponsesRequest) GetExtraParams() map[string]interface{} { + if request == nil || request.ExtraParams == nil { + return make(map[string]interface{}, 0) + } + return request.ExtraParams +} + +// GigaChatResponsesModelOptions contains v2 generation controls. +type GigaChatResponsesModelOptions struct { + Preset *string `json:"preset,omitempty"` + Temperature *float64 `json:"temperature,omitempty"` + TopP *float64 `json:"top_p,omitempty"` + MaxTokens *int `json:"max_tokens,omitempty"` + RepetitionPenalty *float64 `json:"repetition_penalty,omitempty"` + UpdateInterval *float64 `json:"update_interval,omitempty"` + UnnormalizedHistory *bool `json:"unnormalized_history,omitempty"` + TopLogProbs *int `json:"top_logprobs,omitempty"` + Reasoning *GigaChatResponsesReasoning `json:"reasoning,omitempty"` + ResponseFormat *GigaChatResponsesResponseFormat `json:"response_format,omitempty"` + ExtraParams map[string]interface{} `json:"-"` +} + +// GigaChatResponsesReasoning contains GigaChat v2 reasoning controls. +type GigaChatResponsesReasoning struct { + Effort string `json:"effort,omitempty"` +} + +// GigaChatResponsesResponseFormat contains GigaChat v2 structured output controls. +type GigaChatResponsesResponseFormat struct { + Type string `json:"type"` + Schema interface{} `json:"schema,omitempty"` + Strict *bool `json:"strict,omitempty"` + Regex *string `json:"regex,omitempty"` +} + +// GigaChatResponsesMessage is a v2 chat message. +type GigaChatResponsesMessage struct { + Role string `json:"role,omitempty"` + MessageID *string `json:"message_id,omitempty"` + Content []GigaChatResponsesContentPart `json:"content,omitempty"` + ToolsStateID *string `json:"tools_state_id,omitempty"` + ToolStateID *string `json:"tool_state_id,omitempty"` + FunctionCall *GigaChatResponsesFunctionCall `json:"function_call,omitempty"` + FinishReason *string `json:"finish_reason,omitempty"` + ExtraParams map[string]interface{} `json:"-"` +} + +// GigaChatResponsesContentPart is a v2 multipart message content item. +type GigaChatResponsesContentPart struct { + Text *string `json:"text,omitempty"` + Files []GigaChatResponsesContentFile `json:"files,omitempty"` + FunctionCall *GigaChatResponsesFunctionCall `json:"function_call,omitempty"` + FunctionResult *GigaChatResponsesFunctionResult `json:"function_result,omitempty"` + InlineData map[string]interface{} `json:"inline_data,omitempty"` +} + +// GigaChatResponsesContentFile is a v2 file reference. +type GigaChatResponsesContentFile struct { + ID string `json:"id"` + Target *string `json:"target,omitempty"` + MIME *string `json:"mime,omitempty"` +} + +// GigaChatResponsesFunctionCall is a v2 function call content item. +type GigaChatResponsesFunctionCall struct { + Name string `json:"name"` + Arguments interface{} `json:"arguments"` +} + +// GigaChatResponsesFunctionResult is a v2 function result content item. +type GigaChatResponsesFunctionResult struct { + Name string `json:"name"` + Result interface{} `json:"result"` +} + +// GigaChatResponsesToolConfig controls v2 tool invocation policy. +type GigaChatResponsesToolConfig struct { + Mode string `json:"mode,omitempty"` + ToolName *string `json:"tool_name,omitempty"` + FunctionName *string `json:"function_name,omitempty"` +} + +// GigaChatResponsesTool is a v2 tool definition. +type GigaChatResponsesTool struct { + CodeInterpreter map[string]interface{} `json:"code_interpreter,omitempty"` + ImageGenerate map[string]interface{} `json:"image_generate,omitempty"` + WebSearch *GigaChatResponsesWebSearchTool `json:"web_search,omitempty"` + URLContentExtraction map[string]interface{} `json:"url_content_extraction,omitempty"` + Model3DGenerate map[string]interface{} `json:"model_3d_generate,omitempty"` + Functions *GigaChatResponsesFunctionsTool `json:"functions,omitempty"` +} + +func (tool GigaChatResponsesTool) MarshalJSON() ([]byte, error) { + fields := schemas.NewOrderedMap() + if tool.CodeInterpreter != nil { + fields.Set("code_interpreter", tool.CodeInterpreter) + } + if tool.ImageGenerate != nil { + fields.Set("image_generate", tool.ImageGenerate) + } + if tool.WebSearch != nil { + fields.Set("web_search", tool.WebSearch) + } + if tool.URLContentExtraction != nil { + fields.Set("url_content_extraction", tool.URLContentExtraction) + } + if tool.Model3DGenerate != nil { + fields.Set("model_3d_generate", tool.Model3DGenerate) + } + if tool.Functions != nil { + fields.Set("functions", tool.Functions) + } + if fields.Len() == 0 { + return []byte("{}"), nil + } + return json.Marshal(fields) +} + +// GigaChatResponsesWebSearchTool configures GigaChat v2 web search. +type GigaChatResponsesWebSearchTool struct { + Type *string `json:"type,omitempty"` + Indexes []string `json:"indexes,omitempty"` + Flags []string `json:"flags,omitempty"` +} + +// GigaChatResponsesFunctionsTool wraps client-defined function specifications. +type GigaChatResponsesFunctionsTool struct { + Specifications []GigaChatResponsesFunctionSpecification `json:"specifications,omitempty"` +} + +// GigaChatResponsesFunctionSpecification describes a client-defined function. +type GigaChatResponsesFunctionSpecification struct { + Name string `json:"name"` + Description *string `json:"description,omitempty"` + Parameters *schemas.ToolFunctionParameters `json:"parameters"` + FewShotExamples []map[string]interface{} `json:"few_shot_examples,omitempty"` + ReturnParameters map[string]interface{} `json:"return_parameters,omitempty"` +} + +// GigaChatResponsesResponse is the v2 chat completions response body used for Bifrost Responses. +type GigaChatResponsesResponse struct { + ID string `json:"id,omitempty"` + Event *string `json:"event,omitempty"` + Object string `json:"object,omitempty"` + Created int `json:"created,omitempty"` + CreatedAt int `json:"created_at,omitempty"` + Model string `json:"model,omitempty"` + Messages []GigaChatResponsesMessage `json:"messages,omitempty"` + Choices []GigaChatResponsesChoice `json:"choices,omitempty"` + FinishReason *string `json:"finish_reason,omitempty"` + Usage *GigaChatChatUsage `json:"usage,omitempty"` + ThreadID *string `json:"thread_id,omitempty"` + MessageID *string `json:"message_id,omitempty"` + ToolsStateID *string `json:"tools_state_id,omitempty"` + ToolExecution interface{} `json:"tool_execution,omitempty"` + AdditionalData interface{} `json:"additional_data,omitempty"` + SystemFingerprint string `json:"system_fingerprint,omitempty"` + ExtraParams map[string]interface{} `json:"-"` +} + +// GigaChatResponsesChoice is a single v2 completion choice. +type GigaChatResponsesChoice struct { + Index int `json:"index"` + Message *GigaChatResponsesMessage `json:"message,omitempty"` + Delta *GigaChatChatStreamDelta `json:"delta,omitempty"` + FinishReason *string `json:"finish_reason,omitempty"` + LogProbs *schemas.BifrostLogProbs `json:"logprobs,omitempty"` +} + +// # ERROR TYPES + +// GigaChatErrorResponse is the common REST API error shape used by GigaChat. +type GigaChatErrorResponse struct { + Status *int `json:"status,omitempty"` + Code json.RawMessage `json:"code,omitempty"` + Message string `json:"message,omitempty"` + Error string `json:"error,omitempty"` + ErrorDescription string `json:"error_description,omitempty"` +} diff --git a/core/providers/gigachat/utils.go b/core/providers/gigachat/utils.go new file mode 100644 index 00000000000..91b7757d612 --- /dev/null +++ b/core/providers/gigachat/utils.go @@ -0,0 +1,449 @@ +package gigachat + +import ( + "bytes" + "crypto/sha256" + "crypto/tls" + "crypto/x509" + "encoding/hex" + "encoding/json" + "fmt" + "os" + "regexp" + "strings" + "sync" + + "github.com/bytedance/sonic" + providerUtils "github.com/maximhq/bifrost/core/providers/utils" + schemas "github.com/maximhq/bifrost/core/schemas" + "github.com/valyala/fasthttp" +) + +var ( + gigaChatAuthSchemePattern = regexp.MustCompile(`(?i)\b(bearer|basic)\s+[^ \t\r\n"',}]+`) + gigaChatPrivateKeyPattern = regexp.MustCompile(`(?s)-----BEGIN [A-Z ]*PRIVATE KEY-----.*?-----END [A-Z ]*PRIVATE KEY-----`) + gigaChatSensitiveAssignmentPattern = regexp.MustCompile(`(?i)(["']?)\b(authorization|access_token|credentials|username|password|cert_file|key_file|ca_bundle_file|private_key|client_key|client_secret|refresh_token)\b(["']?)(\s*[:=]\s*)("(?:\\.|[^"\\])*"|'(?:\\.|[^'\\])*'|[^ \t\r\n"',}]+)`) + gigaChatUserAssignmentPattern = regexp.MustCompile(`(?i)(["']?)\b(user)\b(["']?)(\s*[:=]\s*)("(?:\\.|[^"\\])*"|'(?:\\.|[^'\\])*'|[^ \t\r\n"',}]+)`) + gigaChatAuthContextTextPattern = regexp.MustCompile(`(?i)(\b(bearer|basic)\s+|\b(authorization|access_token|credentials|username|password|cert_file|key_file|ca_bundle_file|private_key|client_key|client_secret|refresh_token)\b\s*[:=])`) +) + +const ( + gigaChatDefaultBaseURL = "https://gigachat.devices.sberbank.ru/api" + gigaChatDefaultAuthURL = "https://ngw.devices.sberbank.ru:9443/api/v2/oauth" + + gigaChatAPIVersionV1 = "v1" + gigaChatAPIVersionV2 = "v2" + + gigaChatTLSClientCacheAuth = "auth" + gigaChatTLSClientCacheDefault = "default" + gigaChatTLSClientCacheStreaming = "streaming" +) + +type gigaChatTLSClientCache struct { + mu sync.Mutex + clients map[string]*fasthttp.Client +} + +func newGigaChatTLSClientCache() *gigaChatTLSClientCache { + return &gigaChatTLSClientCache{clients: make(map[string]*fasthttp.Client)} +} + +func resolveAuthURL(key schemas.Key) string { + if key.GigaChatKeyConfig != nil { + if authURL := strings.TrimSpace(key.GigaChatKeyConfig.AuthURL); authURL != "" { + return strings.TrimRight(authURL, "/") + } + } + return gigaChatDefaultAuthURL +} + +func resolveBaseURL(key schemas.Key, networkConfig schemas.NetworkConfig) string { + if key.GigaChatKeyConfig != nil { + if baseURL := strings.TrimSpace(key.GigaChatKeyConfig.BaseURL); baseURL != "" { + return strings.TrimRight(baseURL, "/") + } + } + if baseURL := strings.TrimSpace(networkConfig.BaseURL); baseURL != "" { + return strings.TrimRight(baseURL, "/") + } + return gigaChatDefaultBaseURL +} + +func buildGigaChatURL(baseURL string, apiVersion string, path string) string { + resolvedBaseURL := strings.TrimRight(strings.TrimSpace(baseURL), "/") + if resolvedBaseURL == "" { + resolvedBaseURL = gigaChatDefaultBaseURL + } + + version := normalizeGigaChatAPIVersion(apiVersion) + versionedBaseURL := buildGigaChatVersionedBaseURL(resolvedBaseURL, version) + normalizedPath := normalizeGigaChatPath(path, version) + if normalizedPath == "" { + return versionedBaseURL + } + return versionedBaseURL + normalizedPath +} + +func buildGigaChatRequestURL(ctx *schemas.BifrostContext, baseURL string, apiVersion string, defaultPath string, customProviderConfig *schemas.CustomProviderConfig, requestType schemas.RequestType) string { + path, isCompleteURL := providerUtils.GetRequestPath(ctx, defaultPath, customProviderConfig, requestType) + if isCompleteURL { + return path + } + return buildGigaChatURL(baseURL, apiVersion, path) +} + +func normalizeGigaChatAPIVersion(apiVersion string) string { + return strings.Trim(strings.TrimSpace(apiVersion), "/") +} + +func buildGigaChatVersionedBaseURL(baseURL string, apiVersion string) string { + if apiVersion == "" { + return baseURL + } + for _, version := range []string{gigaChatAPIVersionV1, gigaChatAPIVersionV2} { + suffix := "/" + version + if strings.HasSuffix(baseURL, suffix) { + return strings.TrimSuffix(baseURL, suffix) + "/" + apiVersion + } + } + return baseURL + "/" + apiVersion +} + +func normalizeGigaChatPath(path string, apiVersion string) string { + normalizedPath := strings.TrimSpace(path) + if normalizedPath == "" { + return "" + } + normalizedPath = "/" + strings.TrimLeft(normalizedPath, "/") + + for _, version := range []string{apiVersion, gigaChatAPIVersionV1, gigaChatAPIVersionV2} { + if version == "" { + continue + } + versionPrefix := "/" + version + if normalizedPath == versionPrefix { + return "" + } + if strings.HasPrefix(normalizedPath, versionPrefix+"/") { + return strings.TrimPrefix(normalizedPath, versionPrefix) + } + } + + return normalizedPath +} + +func buildGigaChatTLSClient(baseClient *fasthttp.Client, keyConfig *schemas.GigaChatKeyConfig) (*fasthttp.Client, error) { + if keyConfig == nil || !gigaChatKeyConfigHasTLSMaterial(keyConfig) { + return baseClient, nil + } + + client := providerUtils.CloneFastHTTPClientConfig(baseClient) + tlsConfig := client.TLSConfig + if tlsConfig == nil { + tlsConfig = &tls.Config{MinVersion: tls.VersionTLS12} + } else { + tlsConfig = tlsConfig.Clone() + } + + if caBundleFile := strings.TrimSpace(keyConfig.CABundleFile); caBundleFile != "" { + caBundlePEM, err := os.ReadFile(caBundleFile) + if err != nil { + return nil, fmt.Errorf("failed to read gigachat_key_config.ca_bundle_file: %w", err) + } + if tlsConfig.RootCAs == nil { + rootCAs, err := x509.SystemCertPool() + if err != nil || rootCAs == nil { + rootCAs = x509.NewCertPool() + } + tlsConfig.RootCAs = rootCAs + } else { + tlsConfig.RootCAs = tlsConfig.RootCAs.Clone() + } + if !tlsConfig.RootCAs.AppendCertsFromPEM(caBundlePEM) { + return nil, fmt.Errorf("failed to parse gigachat_key_config.ca_bundle_file") + } + } + + hasCertFile := strings.TrimSpace(keyConfig.CertFile) != "" + hasKeyFile := strings.TrimSpace(keyConfig.KeyFile) != "" + if hasCertFile != hasKeyFile { + return nil, fmt.Errorf("gigachat_key_config.cert_file and gigachat_key_config.key_file must be set together") + } + if hasCertFile { + certificate, err := tls.LoadX509KeyPair(keyConfig.CertFile, keyConfig.KeyFile) + if err != nil { + return nil, fmt.Errorf("failed to load gigachat_key_config.cert_file/key_file: %w", err) + } + tlsConfig.Certificates = append(tlsConfig.Certificates, certificate) + } + + client.TLSConfig = tlsConfig + return client, nil +} + +func (provider *GigaChatProvider) getGigaChatTLSClient(baseClient *fasthttp.Client, cacheKind string, keyConfig *schemas.GigaChatKeyConfig) (*fasthttp.Client, error) { + if keyConfig == nil || !gigaChatKeyConfigHasTLSMaterial(keyConfig) { + return baseClient, nil + } + if provider == nil || provider.tlsClientCache == nil { + return buildGigaChatTLSClient(baseClient, keyConfig) + } + + fingerprint := gigaChatTLSConfigFingerprint(keyConfig) + cacheKey := cacheKind + ":" + fingerprint + provider.tlsClientCache.mu.Lock() + client := provider.tlsClientCache.clients[cacheKey] + provider.tlsClientCache.mu.Unlock() + if client != nil { + return client, nil + } + + client, err := buildGigaChatTLSClient(baseClient, keyConfig) + if err != nil { + return nil, err + } + + provider.tlsClientCache.mu.Lock() + defer provider.tlsClientCache.mu.Unlock() + if cached := provider.tlsClientCache.clients[cacheKey]; cached != nil { + return cached, nil + } + provider.tlsClientCache.clients[cacheKey] = client + return client, nil +} + +// gigaChatTLSConfigFingerprint identifies the configured TLS paths without +// reading them, so cache hits never pay certificate file I/O on the request +// path and keep reusing the same connection pool. Replacing material in place +// intentionally takes effect on provider reload rather than on the next request. +func gigaChatTLSConfigFingerprint(keyConfig *schemas.GigaChatKeyConfig) string { + if keyConfig == nil { + return "" + } + hash := sha256.New() + _, _ = hash.Write([]byte("gigachat-tls-config-v1")) + for _, material := range []struct { + field string + path string + }{ + {field: "ca_bundle_file", path: strings.TrimSpace(keyConfig.CABundleFile)}, + {field: "cert_file", path: strings.TrimSpace(keyConfig.CertFile)}, + {field: "key_file", path: strings.TrimSpace(keyConfig.KeyFile)}, + } { + _, _ = hash.Write([]byte{0}) + _, _ = hash.Write([]byte(material.field)) + _, _ = hash.Write([]byte{0}) + _, _ = hash.Write([]byte(material.path)) + } + return hex.EncodeToString(hash.Sum(nil)) +} + +func gigaChatAuthTLSConfigFingerprint(keyConfig *schemas.GigaChatKeyConfig) string { + return gigaChatTLSConfigFingerprint(gigaChatAuthTLSKeyConfig(keyConfig)) +} + +func gigaChatAuthTLSKeyConfig(keyConfig *schemas.GigaChatKeyConfig) *schemas.GigaChatKeyConfig { + if keyConfig == nil { + return nil + } + authKeyConfig := &schemas.GigaChatKeyConfig{ + CABundleFile: strings.TrimSpace(keyConfig.CABundleFile), + } + if !gigaChatKeyConfigHasTLSMaterial(authKeyConfig) { + return nil + } + return authKeyConfig +} + +func gigaChatKeyConfigHasTLSMaterial(keyConfig *schemas.GigaChatKeyConfig) bool { + return strings.TrimSpace(keyConfig.CABundleFile) != "" || + strings.TrimSpace(keyConfig.CertFile) != "" || + strings.TrimSpace(keyConfig.KeyFile) != "" +} + +func enrichGigaChatError(ctx *schemas.BifrostContext, bifrostErr *schemas.BifrostError, requestBody []byte, responseBody []byte, sendBackRawRequest bool, sendBackRawResponse bool) *schemas.BifrostError { + enriched := providerUtils.EnrichError(ctx, bifrostErr, redactGigaChatRawPayload(requestBody), redactGigaChatRawPayload(responseBody), sendBackRawRequest, sendBackRawResponse) + if enriched == nil { + return nil + } + enriched.ExtraFields.RawRequest = redactGigaChatRawValue(enriched.ExtraFields.RawRequest) + enriched.ExtraFields.RawResponse = redactGigaChatRawValue(enriched.ExtraFields.RawResponse) + if enriched.Error != nil { + enriched.Error.Message = redactGigaChatSensitiveText(enriched.Error.Message) + } + return enriched +} + +func redactGigaChatRawPayload(payload []byte) []byte { + if len(payload) == 0 { + return payload + } + var value interface{} + if err := sonic.Unmarshal(payload, &value); err != nil { + return []byte(redactGigaChatSensitiveText(string(payload))) + } + if stringValue, ok := value.(string); ok { + redacted := redactGigaChatSensitiveText(stringValue) + if redacted == stringValue { + return payload + } + redactedPayload, err := sonic.Marshal(redacted) + if err != nil { + return []byte(redacted) + } + return redactedPayload + } + if !redactGigaChatJSONValue(value) { + return payload + } + redacted, err := sonic.Marshal(value) + if err != nil { + return []byte(redactGigaChatSensitiveText(string(payload))) + } + return redacted +} + +func redactGigaChatRawValue(raw interface{}) interface{} { + switch typed := raw.(type) { + case nil: + return nil + case json.RawMessage: + return json.RawMessage(redactGigaChatRawPayload([]byte(typed))) + case []byte: + redacted := redactGigaChatRawPayload(typed) + if json.Valid(redacted) { + return json.RawMessage(redacted) + } + return string(redacted) + case string: + return string(redactGigaChatRawPayload([]byte(typed))) + default: + payload, err := sonic.Marshal(raw) + if err != nil { + return raw + } + redactedPayload := redactGigaChatRawPayload(payload) + if bytes.Equal(payload, redactedPayload) { + return raw + } + var redacted interface{} + if err := sonic.Unmarshal(redactedPayload, &redacted); err != nil { + return string(redactedPayload) + } + return redacted + } +} + +func redactGigaChatJSONValue(value interface{}) bool { + return redactGigaChatJSONValueInContext(value, false) +} + +func redactGigaChatJSONValueInContext(value interface{}, inGigaChatKeyConfig bool) bool { + changed := false + switch typed := value.(type) { + case map[string]interface{}: + authFieldContext := inGigaChatKeyConfig || hasGigaChatAuthSensitiveField(typed) + for key, child := range typed { + childInGigaChatKeyConfig := inGigaChatKeyConfig || strings.EqualFold(strings.TrimSpace(key), "gigachat_key_config") + if isGigaChatSensitiveField(key, authFieldContext) { + typed[key] = "" + changed = true + continue + } + if redactedValue, ok := child.(string); ok { + redacted := redactGigaChatSensitiveText(redactedValue) + if redacted != redactedValue { + typed[key] = redacted + changed = true + } + continue + } + if redactGigaChatJSONValueInContext(child, childInGigaChatKeyConfig) { + changed = true + } + } + case []interface{}: + for index, child := range typed { + if redactedValue, ok := child.(string); ok { + redacted := redactGigaChatSensitiveText(redactedValue) + if redacted != redactedValue { + typed[index] = redacted + changed = true + } + continue + } + if redactGigaChatJSONValueInContext(child, inGigaChatKeyConfig) { + changed = true + } + } + } + return changed +} + +func hasGigaChatAuthSensitiveField(fields map[string]interface{}) bool { + for fieldName := range fields { + if isGigaChatSensitiveField(fieldName, false) { + return true + } + } + return false +} + +func isGigaChatSensitiveField(fieldName string, inGigaChatKeyConfig bool) bool { + switch strings.ToLower(strings.TrimSpace(fieldName)) { + case "authorization", "access_token", "credentials", "username", "password", "cert_file", "key_file", "ca_bundle_file", "private_key", "client_key", "client_secret", "refresh_token": + return true + case "user": + return inGigaChatKeyConfig + default: + return false + } +} + +func redactGigaChatSensitiveText(text string) string { + redacted := text + redacted = gigaChatPrivateKeyPattern.ReplaceAllString(redacted, "") + redacted = gigaChatAuthSchemePattern.ReplaceAllString(redacted, "$1 ") + redacted = redactGigaChatSensitiveAssignments(redacted) + return redacted +} + +func redactGigaChatSensitiveAssignments(text string) string { + redacted := redactGigaChatAssignmentsWithPattern(text, gigaChatSensitiveAssignmentPattern) + if gigaChatAuthContextTextPattern.MatchString(text) { + redacted = redactGigaChatAssignmentsWithPattern(redacted, gigaChatUserAssignmentPattern) + } + return redacted +} + +func redactGigaChatAssignmentsWithPattern(text string, pattern *regexp.Regexp) string { + return pattern.ReplaceAllStringFunc(text, func(match string) string { + parts := pattern.FindStringSubmatch(match) + if len(parts) != 6 { + return "" + } + if parts[1] != parts[3] { + return match + } + + value := "" + if quote := firstGigaChatQuote(parts[5]); quote != "" { + value = quote + value + quote + } + return parts[1] + parts[2] + parts[3] + parts[4] + value + }) +} + +func firstGigaChatQuote(value string) string { + if value == "" { + return "" + } + switch value[0] { + case '"': + return `"` + case '\'': + return `'` + default: + return "" + } +} diff --git a/core/schemas/account.go b/core/schemas/account.go index 0d235754fa9..9fe89d04175 100644 --- a/core/schemas/account.go +++ b/core/schemas/account.go @@ -139,6 +139,7 @@ type Key struct { ReplicateKeyConfig *ReplicateKeyConfig `json:"replicate_key_config,omitempty"` // Replicate-specific key configuration OllamaKeyConfig *OllamaKeyConfig `json:"ollama_key_config,omitempty"` // Ollama-specific key configuration SGLKeyConfig *SGLKeyConfig `json:"sgl_key_config,omitempty"` // SGLang-specific key configuration + GigaChatKeyConfig *GigaChatKeyConfig `json:"gigachat_key_config,omitempty"` // GigaChat-specific key configuration Enabled *bool `json:"enabled,omitempty"` // Whether the key is active (default:true) UseForBatchAPI *bool `json:"use_for_batch_api,omitempty"` // Whether this key can be used for batch API operations (default:false for new keys, migrated keys default to true) UseAnthropicEndpoints *bool `json:"use_anthropic_endpoints,omitempty"` // Whether to use anthropic endpoints for this key @@ -823,6 +824,112 @@ type SGLKeyConfig struct { URL SecretVar `json:"url"` // SGLang server base URL (required, supports env. prefix) } +const ( + // DefaultGigaChatScope is the personal API scope used by GigaChat SDKs and REST examples. + DefaultGigaChatScope = "GIGACHAT_API_PERS" +) + +// GigaChatKeyConfig represents GigaChat-specific authentication and endpoint settings. +type GigaChatKeyConfig struct { + Credentials *SecretVar `json:"credentials,omitempty"` // Authorization key for OAuth token exchange (supports env.* and vault.*) + Scope string `json:"scope,omitempty"` // OAuth scope. Defaults to GIGACHAT_API_PERS. + User *SecretVar `json:"user,omitempty"` // Username for password auth mode (supports env.* and vault.*) + Password *SecretVar `json:"password,omitempty"` // Password for password auth mode (supports env.* and vault.*) + AccessToken *SecretVar `json:"access_token,omitempty"` // Pre-obtained access token (supports env.* and vault.*) + AuthURL string `json:"auth_url,omitempty"` // OAuth token endpoint override + BaseURL string `json:"base_url,omitempty"` // API base URL override for this key + CertFile string `json:"cert_file,omitempty"` // Client certificate file for mTLS + KeyFile string `json:"key_file,omitempty"` // Client private key file for mTLS + CABundleFile string `json:"ca_bundle_file,omitempty"` // CA bundle file for GigaChat TLS roots +} + +// CheckAndSetDefaults applies GigaChat defaults that are safe at config-parse time. +func (config *GigaChatKeyConfig) CheckAndSetDefaults() { + if config == nil { + return + } + if strings.TrimSpace(config.Scope) == "" { + config.Scope = DefaultGigaChatScope + } +} + +// Validate checks static GigaChat auth configuration constraints. +func (config *GigaChatKeyConfig) Validate() error { + if config == nil { + return nil + } + config.CheckAndSetDefaults() + + hasUser := config.User.IsSet() + hasPassword := config.Password.IsSet() + if hasUser != hasPassword { + return fmt.Errorf("gigachat_key_config.user and gigachat_key_config.password must be set together") + } + + hasCertFile := strings.TrimSpace(config.CertFile) != "" + hasKeyFile := strings.TrimSpace(config.KeyFile) != "" + if hasCertFile != hasKeyFile { + return fmt.Errorf("gigachat_key_config.cert_file and gigachat_key_config.key_file must be set together") + } + return nil +} + +// HasAuthMaterial reports whether the config contains a usable bearer auth mode. +func (config *GigaChatKeyConfig) HasAuthMaterial() bool { + if config == nil { + return false + } + if config.AccessToken.IsSet() || config.Credentials.IsSet() { + return true + } + if config.User.IsSet() && config.Password.IsSet() { + return true + } + return false +} + +// HasTLSMaterial reports whether the config contains TLS or mTLS material. +func (config *GigaChatKeyConfig) HasTLSMaterial() bool { + if config == nil { + return false + } + return strings.TrimSpace(config.CertFile) != "" || + strings.TrimSpace(config.KeyFile) != "" || + strings.TrimSpace(config.CABundleFile) != "" +} + +// HasClientCertificateMaterial reports whether the config contains a complete +// client certificate pair for mTLS authentication. +func (config *GigaChatKeyConfig) HasClientCertificateMaterial() bool { + if config == nil { + return false + } + return strings.TrimSpace(config.CertFile) != "" && strings.TrimSpace(config.KeyFile) != "" +} + +// Redacted returns a copy of the GigaChat key config with sensitive fields masked. +func (config *GigaChatKeyConfig) Redacted() *GigaChatKeyConfig { + if config == nil { + return nil + } + redacted := *config + redacted.Credentials = config.Credentials.FullyRedacted() + redacted.User = config.User.FullyRedacted() + redacted.Password = config.Password.FullyRedacted() + redacted.AccessToken = config.AccessToken.FullyRedacted() + redacted.CertFile = redactNonEmptyString(config.CertFile) + redacted.KeyFile = redactNonEmptyString(config.KeyFile) + redacted.CABundleFile = redactNonEmptyString(config.CABundleFile) + return &redacted +} + +func redactNonEmptyString(value string) string { + if strings.TrimSpace(value) == "" { + return value + } + return "" +} + // Account defines the interface for managing provider accounts and their configurations. // It provides methods to access provider-specific settings, API keys, and configurations. type Account interface { diff --git a/core/schemas/bifrost.go b/core/schemas/bifrost.go index cb979add3e0..248fc3129f2 100644 --- a/core/schemas/bifrost.go +++ b/core/schemas/bifrost.go @@ -61,6 +61,7 @@ const ( Cerebras ModelProvider = "cerebras" DeepSeek ModelProvider = "deepseek" Gemini ModelProvider = "gemini" + GigaChat ModelProvider = "gigachat" OpenRouter ModelProvider = "openrouter" Elevenlabs ModelProvider = "elevenlabs" HuggingFace ModelProvider = "huggingface" @@ -96,6 +97,7 @@ var StandardProviders = []ModelProvider{ Cohere, DeepSeek, Gemini, + GigaChat, Groq, Mistral, Ollama, diff --git a/core/utils.go b/core/utils.go index 1d8dd3a91cf..a34c01f55ce 100644 --- a/core/utils.go +++ b/core/utils.go @@ -130,10 +130,10 @@ func providerRequiresKey(customConfig *schemas.CustomProviderConfig) bool { } // CanProviderKeyValueBeEmpty returns true if the given provider allows the API key to be empty. -// Some providers like Vertex and Bedrock have their credentials in additional key configs. +// Some providers like Vertex, Bedrock, Bedrock Mantle, and GigaChat have their credentials in additional key configs. // Ollama and SGL are keyless (API Key is optional) but use per-key server URLs. func CanProviderKeyValueBeEmpty(providerKey schemas.ModelProvider) bool { - return providerKey == schemas.Vertex || providerKey == schemas.Bedrock || providerKey == schemas.BedrockMantle || providerKey == schemas.VLLM || providerKey == schemas.Azure || providerKey == schemas.Ollama || providerKey == schemas.SGL + return providerKey == schemas.Vertex || providerKey == schemas.Bedrock || providerKey == schemas.BedrockMantle || providerKey == schemas.VLLM || providerKey == schemas.Azure || providerKey == schemas.Ollama || providerKey == schemas.SGL || providerKey == schemas.GigaChat } func isKeySkippingAllowed(providerKey schemas.ModelProvider) bool { @@ -212,6 +212,15 @@ func validateKey(providerKey schemas.ModelProvider, key *schemas.Key) error { if key.SGLKeyConfig.URL.GetValue() == "" { return fmt.Errorf("sgl_key_config.url is required") } + case schemas.GigaChat: + if key.GigaChatKeyConfig != nil { + if err := key.GigaChatKeyConfig.Validate(); err != nil { + return err + } + } + if !key.Value.IsSet() && (key.GigaChatKeyConfig == nil || (!key.GigaChatKeyConfig.HasAuthMaterial() && !key.GigaChatKeyConfig.HasClientCertificateMaterial())) { + return fmt.Errorf("gigachat key requires value access token, gigachat_key_config bearer auth material, or gigachat_key_config mTLS client certificate material") + } } return nil } From 8cafef9d58fca482669cc1e70b6bb518a7cdf4f2 Mon Sep 17 00:00:00 2001 From: krakenalt Date: Thu, 20 Aug 2026 13:48:43 +0300 Subject: [PATCH 2/6] [feat]: add GigaChat key configuration Persist and validate GigaChat credentials, TLS paths, and transport schema fields. --- framework/changelog.md | 1 + framework/configstore/clientconfig.go | 12 ++ framework/configstore/migrations.go | 57 +++++++ framework/configstore/rdb.go | 5 + framework/configstore/tables/key.go | 57 +++++++ .../bifrost-http/handlers/provider_keys.go | 55 +++++++ transports/bifrost-http/lib/config.go | 3 + transports/config.schema.json | 145 ++++++++++++++++++ 8 files changed, 335 insertions(+) diff --git a/framework/changelog.md b/framework/changelog.md index 515b3eb56a7..e06c6821523 100644 --- a/framework/changelog.md +++ b/framework/changelog.md @@ -1,3 +1,4 @@ +- **GigaChat Key Reconciliation**: Include GigaChat authentication and TLS configuration in key reconciliation hashes. [@krakenalt](https://github.com/krakenalt) - feat: add `cost_per_request` flat-fee pricing field across DB, cost engine, overrides and docs (#6079) - feat(modelcatalog): resolve pricing overrides for catalog rows (#6055) - feat: add `use_idp_credentials` to token-exchange config (#6068) diff --git a/framework/configstore/clientconfig.go b/framework/configstore/clientconfig.go index 4453887728e..50c1030145a 100644 --- a/framework/configstore/clientconfig.go +++ b/framework/configstore/clientconfig.go @@ -673,6 +673,10 @@ func (p *ProviderConfig) Redacted() *ProviderConfig { sglConfig.URL = *key.SGLKeyConfig.URL.Redacted() redactedConfig.Keys[i].SGLKeyConfig = sglConfig } + + if key.GigaChatKeyConfig != nil { + redactedConfig.Keys[i].GigaChatKeyConfig = key.GigaChatKeyConfig.Redacted() + } } return &redactedConfig } @@ -863,6 +867,14 @@ func GenerateKeyHash(key schemas.Key) (string, error) { } hash.Write(data) } + // Hash GigaChatKeyConfig + if key.GigaChatKeyConfig != nil { + data, err := sonic.Marshal(key.GigaChatKeyConfig) + if err != nil { + return "", err + } + hash.Write(data) + } // Hash Enabled (nil = false, only true produces different hash) if key.Enabled != nil && *key.Enabled { hash.Write([]byte("enabled:true")) diff --git a/framework/configstore/migrations.go b/framework/configstore/migrations.go index 0e43aeb95be..1df39af82e9 100644 --- a/framework/configstore/migrations.go +++ b/framework/configstore/migrations.go @@ -368,6 +368,7 @@ var configstoreMigrationSteps = []migrationStep{ {IDs: []string{"add_prompt_variables_columns"}, run: migrationAddPromptVariablesColumns}, {IDs: []string{"add_model_capability_columns"}, run: migrationAddModelCapabilityColumns}, {IDs: []string{"add_ollama_sgl_config_columns"}, run: migrationAddOllamaSGLConfigColumns}, + {IDs: []string{"add_gigachat_key_config_column"}, run: migrationAddGigaChatKeyConfigColumn}, {IDs: []string{"add_multi_budget_tables"}, run: migrationAddMultiBudgetTables}, {IDs: []string{"add_per_user_oauth_tables"}, run: migrationAddPerUserOAuthTables}, {IDs: []string{"add_mcp_client_discovered_tools_columns"}, run: migrationAddMCPClientDiscoveredToolsColumns}, @@ -7406,6 +7407,62 @@ func migrationAddOllamaSGLConfigColumns(ctx context.Context, db *gorm.DB, logger return nil } +// migrationAddGigaChatKeyConfigColumn adds GigaChat key auth config storage. +func migrationAddGigaChatKeyConfigColumn(ctx context.Context, db *gorm.DB, logger schemas.Logger) error { + migrationName := "add_gigachat_key_config_column" + logger.Info("[configstore] starting migration %s", migrationName) + defer logger.Info("[configstore] finished migration %s", migrationName) + m := migrator.New(db, migrator.DefaultOptions, []*migrator.Migration{{ + ID: migrationName, + Migrate: func(tx *gorm.DB) error { + tx = tx.WithContext(ctx) + migrator := tx.Migrator() + if !migrator.HasColumn(&tables.TableKey{}, "gigachat_key_config_json") { + logger.Info("[configstore] %s: adding column gigachat_key_config_json to TableKey", migrationName) + if err := migrator.AddColumn(&tables.TableKey{}, "gigachat_key_config_json"); err != nil { + return err + } + } + + // GigaChatKeyConfig is part of GenerateKeyHash. Refresh existing GigaChat + // rows so reconciliation compares config.json against the complete stored + // key configuration after this column becomes available. + var affectedKeys []tables.TableKey + if err := tx.Where("provider = ?", string(schemas.GigaChat)).Find(&affectedKeys).Error; err != nil { + return fmt.Errorf("failed to fetch GigaChat keys for hash recomputation: %w", err) + } + logger.Info("[configstore] %s: processing %d affectedKeys", migrationName, len(affectedKeys)) + for _, key := range affectedKeys { + hash, err := GenerateKeyHash(schemaKeyFromTableKey(key)) + if err != nil { + return fmt.Errorf("failed to generate hash for GigaChat key %s: %w", key.Name, err) + } + if err := tx.Model(&tables.TableKey{}). + Where("id = ?", key.ID). + Update("config_hash", hash).Error; err != nil { + return fmt.Errorf("failed to update config_hash for GigaChat key %s: %w", key.Name, err) + } + } + return nil + }, + Rollback: func(tx *gorm.DB) error { + tx = tx.WithContext(ctx) + migrator := tx.Migrator() + if migrator.HasColumn(&tables.TableKey{}, "gigachat_key_config_json") { + logger.Info("[configstore] %s: dropping column gigachat_key_config_json from TableKey", migrationName) + if err := migrator.DropColumn(&tables.TableKey{}, "gigachat_key_config_json"); err != nil { + return err + } + } + return nil + }, + }}) + if err := m.Migrate(); err != nil { + return fmt.Errorf("error while running gigachat key config column migration: %s", err.Error()) + } + return nil +} + // migrationAddMultiBudgetTables creates junction tables for multi-budget support and backfills existing data. func migrationAddMultiBudgetTables(ctx context.Context, db *gorm.DB, logger schemas.Logger) error { migrationName := "add_multi_budget_tables" diff --git a/framework/configstore/rdb.go b/framework/configstore/rdb.go index 275847f9051..e53f0c7daf8 100644 --- a/framework/configstore/rdb.go +++ b/framework/configstore/rdb.go @@ -163,6 +163,7 @@ func schemaKeyFromTableKey(dbKey tables.TableKey) schemas.Key { ReplicateKeyConfig: dbKey.ReplicateKeyConfig, OllamaKeyConfig: dbKey.OllamaKeyConfig, SGLKeyConfig: dbKey.SGLKeyConfig, + GigaChatKeyConfig: dbKey.GigaChatKeyConfig, ConfigHash: dbKey.ConfigHash, Status: schemas.KeyStatusType(dbKey.Status), Description: dbKey.Description, @@ -192,6 +193,7 @@ func tableKeyFromSchemaKey(provider tables.TableProvider, key schemas.Key) (tabl ReplicateKeyConfig: key.ReplicateKeyConfig, OllamaKeyConfig: key.OllamaKeyConfig, SGLKeyConfig: key.SGLKeyConfig, + GigaChatKeyConfig: key.GigaChatKeyConfig, ConfigHash: key.ConfigHash, Status: string(key.Status), Description: key.Description, @@ -741,6 +743,7 @@ func (s *RDBConfigStore) UpdateProvidersConfig(ctx context.Context, providers ma ReplicateKeyConfig: key.ReplicateKeyConfig, OllamaKeyConfig: key.OllamaKeyConfig, SGLKeyConfig: key.SGLKeyConfig, + GigaChatKeyConfig: key.GigaChatKeyConfig, ConfigHash: keyHash, Status: string(key.Status), Description: key.Description, @@ -982,6 +985,7 @@ func (s *RDBConfigStore) UpdateProvider(ctx context.Context, provider schemas.Mo ReplicateKeyConfig: key.ReplicateKeyConfig, OllamaKeyConfig: key.OllamaKeyConfig, SGLKeyConfig: key.SGLKeyConfig, + GigaChatKeyConfig: key.GigaChatKeyConfig, ConfigHash: keyHash, Status: string(key.Status), Description: key.Description, @@ -1135,6 +1139,7 @@ func (s *RDBConfigStore) AddProvider(ctx context.Context, provider schemas.Model ReplicateKeyConfig: key.ReplicateKeyConfig, OllamaKeyConfig: key.OllamaKeyConfig, SGLKeyConfig: key.SGLKeyConfig, + GigaChatKeyConfig: key.GigaChatKeyConfig, ConfigHash: key.ConfigHash, Status: string(key.Status), Description: key.Description, diff --git a/framework/configstore/tables/key.go b/framework/configstore/tables/key.go index 9754ac2b95c..7280a9cfaa7 100644 --- a/framework/configstore/tables/key.go +++ b/framework/configstore/tables/key.go @@ -84,6 +84,9 @@ type TableKey struct { // SGL config fields (embedded) SGLUrl *schemas.SecretVar `gorm:"type:text" json:"sgl_url,omitempty"` + // GigaChat config fields (serialized because the auth surface spans multiple optional modes) + GigaChatKeyConfigJSON *string `gorm:"column:gigachat_key_config_json;type:text" json:"-"` + // Batch API configuration UseForBatchAPI *bool `gorm:"default:false" json:"use_for_batch_api,omitempty"` // Whether this key can be used for batch API operations @@ -108,6 +111,7 @@ type TableKey struct { ReplicateKeyConfig *schemas.ReplicateKeyConfig `gorm:"-" json:"replicate_key_config,omitempty"` OllamaKeyConfig *schemas.OllamaKeyConfig `gorm:"-" json:"ollama_key_config,omitempty"` SGLKeyConfig *schemas.SGLKeyConfig `gorm:"-" json:"sgl_key_config,omitempty"` + GigaChatKeyConfig *schemas.GigaChatKeyConfig `gorm:"-" json:"gigachat_key_config,omitempty"` } // TableName sets the table name for each model @@ -449,6 +453,41 @@ func (k *TableKey) BeforeSave(tx *gorm.DB) error { k.SGLUrl = nil } + if k.GigaChatKeyConfig != nil { + cfg := *k.GigaChatKeyConfig + if cfg.Credentials != nil { + v := *cfg.Credentials + cfg.Credentials = &v + } + if cfg.User != nil { + v := *cfg.User + cfg.User = &v + } + if cfg.Password != nil { + v := *cfg.Password + cfg.Password = &v + } + if cfg.AccessToken != nil { + v := *cfg.AccessToken + cfg.AccessToken = &v + } + cfg.CheckAndSetDefaults() + if schemas.VaultStoreWriteEnabled() { + base := schemas.VaultBasePath(tx.Statement.Table, k.VaultPathKey()) + "/gigachat_key_config" + if err := schemas.StoreOwnedVaultSecretVars(tx.Statement.Context, base, &cfg); err != nil { + return fmt.Errorf("failed to store gigachat key secrets to vault: %w", err) + } + } + data, err := sonic.Marshal(&cfg) + if err != nil { + return err + } + s := string(data) + k.GigaChatKeyConfigJSON = &s + } else { + k.GigaChatKeyConfigJSON = nil + } + // Store plaintext SecretVar columns into the vault and rewrite them to vault refs. // This must run after the columns are populated (above) and before encryption (below): // encryptSecretVar skips fields that are already vault refs, so vault-owned secrets are @@ -573,6 +612,10 @@ func (k *TableKey) BeforeSave(tx *gorm.DB) error { if err := encryptSecretVarPtr(&k.SGLUrl); err != nil { return fmt.Errorf("failed to encrypt sgl url: %w", err) } + // GigaChat + if err := encryptString(k.GigaChatKeyConfigJSON); err != nil { + return fmt.Errorf("failed to encrypt gigachat key config: %w", err) + } k.EncryptionStatus = EncryptionStatusEncrypted } return nil @@ -694,6 +737,10 @@ func (k *TableKey) AfterFind(tx *gorm.DB) error { if err := decryptSecretVarPtr(&k.SGLUrl); err != nil { return fmt.Errorf("failed to decrypt sgl url: %w", err) } + // GigaChat + if err := decryptString(k.GigaChatKeyConfigJSON); err != nil { + return fmt.Errorf("failed to decrypt gigachat key config: %w", err) + } } if k.ModelsJSON != "" { @@ -873,6 +920,16 @@ func (k *TableKey) AfterFind(tx *gorm.DB) error { } else { k.SGLKeyConfig = nil } + if k.GigaChatKeyConfigJSON != nil && *k.GigaChatKeyConfigJSON != "" { + var config schemas.GigaChatKeyConfig + if err := sonic.Unmarshal([]byte(*k.GigaChatKeyConfigJSON), &config); err != nil { + return err + } + config.CheckAndSetDefaults() + k.GigaChatKeyConfig = &config + } else { + k.GigaChatKeyConfig = nil + } return nil } diff --git a/transports/bifrost-http/handlers/provider_keys.go b/transports/bifrost-http/handlers/provider_keys.go index 003da12f080..8349d004257 100644 --- a/transports/bifrost-http/handlers/provider_keys.go +++ b/transports/bifrost-http/handlers/provider_keys.go @@ -4,6 +4,7 @@ import ( "errors" "fmt" "net/url" + "strings" "github.com/bytedance/sonic" "github.com/google/uuid" @@ -563,6 +564,51 @@ func (h *ProviderHandler) mergeUpdatedKey(oldRawKey, updateKey schemas.Key) (sch } } + if mergedKey.GigaChatKeyConfig != nil { + var credentials, user, password, accessToken *schemas.SecretVar + var certFile, keyFile, caBundleFile string + if oldRawKey.GigaChatKeyConfig != nil { + credentials = oldRawKey.GigaChatKeyConfig.Credentials + user = oldRawKey.GigaChatKeyConfig.User + password = oldRawKey.GigaChatKeyConfig.Password + accessToken = oldRawKey.GigaChatKeyConfig.AccessToken + certFile = oldRawKey.GigaChatKeyConfig.CertFile + keyFile = oldRawKey.GigaChatKeyConfig.KeyFile + caBundleFile = oldRawKey.GigaChatKeyConfig.CABundleFile + } + for _, item := range []struct { + incoming *schemas.SecretVar + stored *schemas.SecretVar + field string + }{ + {mergedKey.GigaChatKeyConfig.Credentials, credentials, "gigachat_key_config.credentials"}, + {mergedKey.GigaChatKeyConfig.User, user, "gigachat_key_config.user"}, + {mergedKey.GigaChatKeyConfig.Password, password, "gigachat_key_config.password"}, + {mergedKey.GigaChatKeyConfig.AccessToken, accessToken, "gigachat_key_config.access_token"}, + } { + if err := preserve(item.incoming, item.stored, item.field); err != nil { + return schemas.Key{}, err + } + } + for _, item := range []struct { + incoming *string + stored string + field string + }{ + {&mergedKey.GigaChatKeyConfig.CertFile, certFile, "gigachat_key_config.cert_file"}, + {&mergedKey.GigaChatKeyConfig.KeyFile, keyFile, "gigachat_key_config.key_file"}, + {&mergedKey.GigaChatKeyConfig.CABundleFile, caBundleFile, "gigachat_key_config.ca_bundle_file"}, + } { + if *item.incoming != "" { + continue + } + if strings.TrimSpace(item.stored) == "" { + return schemas.Key{}, fmt.Errorf("masked preview cannot be used for %s without a stored value", item.field) + } + *item.incoming = item.stored + } + } + mergedKey.ConfigHash = oldRawKey.ConfigHash mergedKey.Status = oldRawKey.Status @@ -602,6 +648,15 @@ func validateProviderKeyURL(provider schemas.ModelProvider, key schemas.Key) err if key.SGLKeyConfig == nil || !key.SGLKeyConfig.URL.IsSet() { return fmt.Errorf("sgl_key_config.url is required for SGL keys") } + case schemas.GigaChat: + if key.GigaChatKeyConfig != nil { + if err := key.GigaChatKeyConfig.Validate(); err != nil { + return err + } + } + if !key.Value.IsSet() && (key.GigaChatKeyConfig == nil || (!key.GigaChatKeyConfig.HasAuthMaterial() && !key.GigaChatKeyConfig.HasClientCertificateMaterial())) { + return fmt.Errorf("gigachat key requires value access token, gigachat_key_config bearer auth material, or gigachat_key_config mTLS client certificate material") + } case schemas.Azure: if key.AzureKeyConfig == nil || !key.AzureKeyConfig.Endpoint.IsSet() { return fmt.Errorf("azure_key_config.endpoint is required for Azure keys") diff --git a/transports/bifrost-http/lib/config.go b/transports/bifrost-http/lib/config.go index 417b9174652..5b948b43354 100644 --- a/transports/bifrost-http/lib/config.go +++ b/transports/bifrost-http/lib/config.go @@ -6524,6 +6524,9 @@ func (c *Config) GetAllKeys() ([]configstoreTables.TableKey, error) { cfg.URL = *cfg.URL.Redacted() configStoreKey.SGLKeyConfig = &cfg } + if key.GigaChatKeyConfig != nil { + configStoreKey.GigaChatKeyConfig = key.GigaChatKeyConfig.Redacted() + } keys = append(keys, configStoreKey) } } diff --git a/transports/config.schema.json b/transports/config.schema.json index 2344a3b86f8..7145bd43e0e 100644 --- a/transports/config.schema.json +++ b/transports/config.schema.json @@ -409,6 +409,9 @@ "gemini": { "$ref": "#/$defs/provider" }, + "gigachat": { + "$ref": "#/$defs/provider_with_gigachat_config" + }, "opencode-go": { "$ref": "#/$defs/provider" }, @@ -2524,6 +2527,7 @@ "openai", "anthropic", "gemini", + "gigachat", "bedrock", "bedrock_mantle", "azure", @@ -4518,6 +4522,108 @@ } ] }, + "gigachat_key": { + "allOf": [ + { + "$ref": "#/$defs/base_key" + }, + { + "type": "object", + "properties": { + "gigachat_key_config": { + "type": "object", + "properties": { + "credentials": { + "type": "string", + "pattern": "\\S", + "description": "GigaChat authorization key for OAuth token exchange (can use env. prefix)" + }, + "scope": { + "type": "string", + "description": "GigaChat OAuth scope, defaults to GIGACHAT_API_PERS" + }, + "user": { + "type": "string", + "pattern": "\\S", + "description": "GigaChat username for password auth mode (can use env. prefix)" + }, + "password": { + "type": "string", + "pattern": "\\S", + "description": "GigaChat password for password auth mode (can use env. prefix)" + }, + "access_token": { + "type": "string", + "pattern": "\\S", + "description": "Pre-obtained GigaChat bearer access token (can use env. prefix)" + }, + "auth_url": { + "type": "string", + "description": "GigaChat OAuth token endpoint override" + }, + "base_url": { + "type": "string", + "description": "GigaChat API base URL override for this key" + }, + "cert_file": { + "type": "string", + "pattern": "\\S", + "description": "Client certificate file for GigaChat mTLS. Complete client certificate material can be used with or without bearer auth." + }, + "key_file": { + "type": "string", + "pattern": "\\S", + "description": "Client private key file for GigaChat mTLS. Must be configured with cert_file." + }, + "ca_bundle_file": { + "type": "string", + "description": "CA bundle file for GigaChat TLS roots." + } + }, + "dependentRequired": { + "user": ["password"], + "password": ["user"], + "cert_file": ["key_file"], + "key_file": ["cert_file"] + }, + "additionalProperties": false + } + }, + "anyOf": [ + { + "required": ["value"], + "properties": { + "value": { + "type": "string", + "pattern": "\\S" + } + } + }, + { + "required": ["gigachat_key_config"], + "properties": { + "gigachat_key_config": { + "anyOf": [ + { + "required": ["credentials"] + }, + { + "required": ["access_token"] + }, + { + "required": ["user", "password"] + }, + { + "required": ["cert_file", "key_file"] + } + ] + } + } + } + ] + } + ] + }, "deepseek_key": { "allOf": [ { @@ -4992,6 +5098,45 @@ }, "additionalProperties": false }, + "provider_with_gigachat_config": { + "type": "object", + "properties": { + "keys": { + "type": "array", + "items": { + "$ref": "#/$defs/gigachat_key" + }, + "minItems": 1, + "description": "GigaChat auth keys and per-key endpoint overrides" + }, + "network_config": { + "$ref": "#/$defs/network_config" + }, + "concurrency_and_buffer_size": { + "$ref": "#/$defs/concurrency_and_buffer_size" + }, + "proxy_config": { + "$ref": "#/$defs/proxy_config" + }, + "send_back_raw_request": { + "type": "boolean", + "description": "Include raw request in BifrostResponse (default: false)" + }, + "send_back_raw_response": { + "type": "boolean", + "description": "Include raw response in BifrostResponse (default: false)" + }, + "store_raw_request_response": { + "type": "boolean", + "description": "Capture raw request/response for internal logging only; strip from API responses returned to clients (default: false)" + }, + "custom_provider_config": { + "$ref": "#/$defs/custom_provider_config" + } + }, + "required": ["keys"], + "additionalProperties": false + }, "provider_with_deepseek_config": { "type": "object", "properties": { From c926f0ea2757b7c08555784a5c813e49149d2d54 Mon Sep 17 00:00:00 2001 From: krakenalt Date: Thu, 20 Aug 2026 13:49:06 +0300 Subject: [PATCH 3/6] [test]: add GigaChat provider coverage Cover provider conversions, auth, files, batches, config persistence, and integration wiring. --- core/internal/llmtests/account.go | 110 + core/providers/gigachat/auth_test.go | 1453 +++++++++ core/providers/gigachat/batch_test.go | 951 ++++++ core/providers/gigachat/chat_test.go | 1870 ++++++++++++ core/providers/gigachat/count_tokens_test.go | 440 +++ core/providers/gigachat/embedding_test.go | 358 +++ core/providers/gigachat/errors_test.go | 254 ++ core/providers/gigachat/files_test.go | 739 +++++ .../gigachat/gigachat_comprehensive_test.go | 60 + .../gigachat/gigachat_integration_test.go | 751 +++++ core/providers/gigachat/gigachat_test.go | 533 ++++ core/providers/gigachat/key_config_test.go | 209 ++ core/providers/gigachat/models_test.go | 296 ++ core/providers/gigachat/responses_test.go | 2720 +++++++++++++++++ core/providers/gigachat/tools_test.go | 1104 +++++++ core/providers/gigachat/types_test.go | 94 + core/providers/gigachat/utils_test.go | 256 ++ .../clientconfig_redaction_test.go | 17 + framework/configstore/migrations_test.go | 120 + .../configstore/tables/encryption_test.go | 33 +- .../handlers/provider_keys_test.go | 43 + .../bifrost-http/handlers/providers_test.go | 70 + transports/bifrost-http/lib/validator_test.go | 221 ++ 23 files changed, 12701 insertions(+), 1 deletion(-) create mode 100644 core/providers/gigachat/auth_test.go create mode 100644 core/providers/gigachat/batch_test.go create mode 100644 core/providers/gigachat/chat_test.go create mode 100644 core/providers/gigachat/count_tokens_test.go create mode 100644 core/providers/gigachat/embedding_test.go create mode 100644 core/providers/gigachat/errors_test.go create mode 100644 core/providers/gigachat/files_test.go create mode 100644 core/providers/gigachat/gigachat_comprehensive_test.go create mode 100644 core/providers/gigachat/gigachat_integration_test.go create mode 100644 core/providers/gigachat/gigachat_test.go create mode 100644 core/providers/gigachat/key_config_test.go create mode 100644 core/providers/gigachat/models_test.go create mode 100644 core/providers/gigachat/responses_test.go create mode 100644 core/providers/gigachat/tools_test.go create mode 100644 core/providers/gigachat/types_test.go create mode 100644 core/providers/gigachat/utils_test.go diff --git a/core/internal/llmtests/account.go b/core/internal/llmtests/account.go index 1a975eac566..1b28cf583e9 100644 --- a/core/internal/llmtests/account.go +++ b/core/internal/llmtests/account.go @@ -183,6 +183,7 @@ func (account *ComprehensiveTestAccount) GetConfiguredProviders() ([]schemas.Mod schemas.Cerebras, schemas.DeepSeek, schemas.Gemini, + schemas.GigaChat, schemas.OpenRouter, schemas.HuggingFace, schemas.Nebius, @@ -220,6 +221,37 @@ func replicateProviderTestKeys() []schemas.Key { } } +func gigaChatProviderTestKey() schemas.Key { + keyConfig := &schemas.GigaChatKeyConfig{ + Scope: getEnvWithDefault("GIGACHAT_SCOPE", schemas.DefaultGigaChatScope), + AuthURL: os.Getenv("GIGACHAT_AUTH_URL"), + BaseURL: os.Getenv("GIGACHAT_BASE_URL"), + CABundleFile: os.Getenv("GIGACHAT_CA_BUNDLE_FILE"), + } + if certFile, keyFile := os.Getenv("GIGACHAT_CERT_FILE"), os.Getenv("GIGACHAT_KEY_FILE"); certFile != "" && keyFile != "" { + keyConfig.CertFile = certFile + keyConfig.KeyFile = keyFile + } + + switch { + case os.Getenv("GIGACHAT_ACCESS_TOKEN") != "": + keyConfig.AccessToken = schemas.NewSecretVar("env.GIGACHAT_ACCESS_TOKEN") + case os.Getenv("GIGACHAT_USER") != "" && os.Getenv("GIGACHAT_PASSWORD") != "" && os.Getenv("GIGACHAT_BASE_URL") != "": + keyConfig.User = schemas.NewSecretVar("env.GIGACHAT_USER") + keyConfig.Password = schemas.NewSecretVar("env.GIGACHAT_PASSWORD") + default: + keyConfig.Credentials = schemas.NewSecretVar("env.GIGACHAT_CREDENTIALS") + } + + return schemas.Key{ + Name: "gigachat-inference", + Models: []string{"*"}, + Weight: 1.0, + UseForBatchAPI: bifrost.Ptr(true), + GigaChatKeyConfig: keyConfig, + } +} + // GetKeysForProvider returns the API keys and associated models for a given provider. func (account *ComprehensiveTestAccount) GetKeysForProvider(ctx context.Context, providerKey schemas.ModelProvider) ([]schemas.Key, error) { switch providerKey { @@ -503,6 +535,10 @@ func (account *ComprehensiveTestAccount) GetKeysForProvider(ctx context.Context, UseForBatchAPI: bifrost.Ptr(true), }, }, nil + case schemas.GigaChat: + return []schemas.Key{ + gigaChatProviderTestKey(), + }, nil case schemas.OpenRouter: return []schemas.Key{ { @@ -912,6 +948,20 @@ func (account *ComprehensiveTestAccount) GetConfigForProvider(providerKey schema BufferSize: 20, }, }, nil + case schemas.GigaChat: + return &schemas.ProviderConfig{ + NetworkConfig: schemas.NetworkConfig{ + BaseURL: os.Getenv("GIGACHAT_BASE_URL"), + DefaultRequestTimeoutInSeconds: 120, + MaxRetries: 10, + RetryBackoffInitial: 1 * time.Second, + RetryBackoffMax: 20 * time.Second, + }, + ConcurrencyAndBufferSize: schemas.ConcurrencyAndBufferSize{ + Concurrency: Concurrency, + BufferSize: 10, + }, + }, nil case schemas.OpenRouter: return &schemas.ProviderConfig{ NetworkConfig: schemas.NetworkConfig{ @@ -1021,6 +1071,65 @@ func (account *ComprehensiveTestAccount) GetConfigForProvider(providerKey schema } } +// GigaChatComprehensiveTestConfig returns the provider checklist scenarios for GigaChat. +func GigaChatComprehensiveTestConfig() ComprehensiveTestConfig { + return ComprehensiveTestConfig{ + Provider: schemas.GigaChat, + ChatModel: getEnvWithDefault("GIGACHAT_CHAT_MODEL", "GigaChat-2"), + TextModel: "", + EmbeddingModel: getEnvWithDefault("GIGACHAT_EMBEDDING_MODEL", "Embeddings"), + Scenarios: TestScenarios{ + TextCompletion: false, + TextCompletionStream: false, + SimpleChat: true, + CompletionStream: true, + MultiTurnConversation: true, + ToolCalls: true, + ToolCallsStreaming: true, + MultipleToolCalls: false, + MultipleToolCallsStreaming: false, + End2EndToolCalling: true, + AutomaticFunctionCall: true, + ImageURL: false, + ImageBase64: false, + MultipleImages: false, + FileBase64: false, + FileURL: false, + CompleteEnd2End: true, + SpeechSynthesis: false, + SpeechSynthesisStream: false, + Transcription: false, + TranscriptionStream: false, + Embedding: true, + Reasoning: false, // Partial passthrough only; the generic suite sends unsupported Responses reasoning fields. + ListModels: true, + ImageGeneration: false, + ImageGenerationStream: false, + ImageEdit: false, + ImageEditStream: false, + ImageVariation: false, + ImageVariationStream: false, + BatchCreate: true, + BatchList: true, + BatchRetrieve: true, + BatchCancel: false, + BatchResults: true, + FileUpload: true, + FileList: true, + FileRetrieve: true, + FileDelete: true, + FileContent: true, + FileBatchInput: true, + CountTokens: true, + StructuredOutputs: true, + WebSearchTool: false, + PassthroughAPI: false, + WebSocketResponses: false, + Realtime: false, + }, + } +} + // AllProviderConfigs contains test configurations for all providers var AllProviderConfigs = []ComprehensiveTestConfig{ { @@ -1542,6 +1651,7 @@ var AllProviderConfigs = []ComprehensiveTestConfig{ {Provider: schemas.OpenAI, Model: "gpt-4o-mini"}, }, }, + GigaChatComprehensiveTestConfig(), { Provider: schemas.OpenRouter, ChatModel: "openai/gpt-4o", diff --git a/core/providers/gigachat/auth_test.go b/core/providers/gigachat/auth_test.go new file mode 100644 index 00000000000..6742ce15d75 --- /dev/null +++ b/core/providers/gigachat/auth_test.go @@ -0,0 +1,1453 @@ +package gigachat + +import ( + "context" + "crypto/tls" + "crypto/x509" + "encoding/base64" + "encoding/pem" + "io" + "net" + "net/http" + "net/http/httptest" + "net/url" + "os" + "strconv" + "strings" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/google/uuid" + + "github.com/maximhq/bifrost/core/schemas" +) + +func TestGigaChatOAuthTokenClient(t *testing.T) { + t.Parallel() + + t.Run("RequestShapeAndDefaultScope", testGigaChatOAuthRequestShapeAndDefaultScope) + t.Run("ParsesMillisecondsExpiresAt", testGigaChatOAuthParsesMillisecondsExpiresAt) + t.Run("CachesTokenBeforeLeeway", testGigaChatOAuthCachesTokenBeforeLeeway) + t.Run("CacheIncludesCABundle", testGigaChatOAuthCacheIncludesCABundle) + t.Run("RefreshesTokenInsideLeeway", testGigaChatOAuthRefreshesTokenInsideLeeway) + t.Run("IgnoresClientCertificate", testGigaChatOAuthIgnoresClientCertificate) + t.Run("HandlesProviderErrors", testGigaChatOAuthHandlesProviderErrors) + t.Run("HandlesMalformedResponses", testGigaChatOAuthHandlesMalformedResponses) + t.Run("MissingCredentials", testGigaChatOAuthMissingCredentials) + t.Run("ContextCancellation", testGigaChatOAuthContextCancellation) +} + +func TestGigaChatPasswordTokenClient(t *testing.T) { + t.Parallel() + + t.Run("RequestShape", testGigaChatPasswordRequestShape) + t.Run("ParsesSecondsExpiresAt", testGigaChatPasswordParsesSecondsExpiresAt) + t.Run("CachesTokenBeforeLeeway", testGigaChatPasswordCachesTokenBeforeLeeway) + t.Run("CacheIncludesCABundle", testGigaChatPasswordCacheIncludesCABundle) + t.Run("RefreshesTokenInsideLeeway", testGigaChatPasswordRefreshesTokenInsideLeeway) + t.Run("IgnoresClientCertificate", testGigaChatPasswordIgnoresClientCertificate) + t.Run("RejectsExpiredToken", testGigaChatPasswordRejectsExpiredToken) + t.Run("HandlesProviderErrors", testGigaChatPasswordHandlesProviderErrors) + t.Run("HandlesMalformedResponses", testGigaChatPasswordHandlesMalformedResponses) + t.Run("MissingUserPassword", testGigaChatPasswordMissingUserPassword) + t.Run("AuthPriority", testGigaChatAuthPriority) +} + +func TestParseGigaChatExpiresAt(t *testing.T) { + t.Parallel() + + seconds := int64(1_700_001_800) + if got := parseGigaChatExpiresAt(seconds); !got.Equal(time.Unix(seconds, 0)) { + t.Fatalf("seconds expiry mismatch: got %s, want %s", got, time.Unix(seconds, 0)) + } + + milliseconds := seconds * 1000 + if got := parseGigaChatExpiresAt(milliseconds); !got.Equal(time.UnixMilli(milliseconds)) { + t.Fatalf("milliseconds expiry mismatch: got %s, want %s", got, time.UnixMilli(milliseconds)) + } +} + +func TestGigaChatTokenCache(t *testing.T) { + t.Parallel() + + t.Run("PrunesExpiredEntriesAndKeepsReusableEntries", testGigaChatTokenCachePrunesExpiredEntriesAndKeepsReusableEntries) + t.Run("ConcurrentAccess", testGigaChatTokenCacheConcurrentAccess) +} + +func TestGigaChatTLSCacheKeys(t *testing.T) { + t.Parallel() + + t.Run("TokenCacheKeysDoNotReadCABundleFiles", testGigaChatTokenCacheKeysDoNotReadCABundleFiles) +} + +func testGigaChatTokenCachePrunesExpiredEntriesAndKeepsReusableEntries(t *testing.T) { + t.Parallel() + + now := time.Unix(1_700_000_000, 0) + cache := newGigaChatTokenCache(func() time.Time { return now }) + cache.entries["expired"] = &gigaChatTokenCacheEntry{ + token: gigaChatCachedToken{ + accessToken: "expired-token", + expiresAt: now.Add(-time.Second), + }, + } + cache.entries["valid"] = &gigaChatTokenCacheEntry{ + token: gigaChatCachedToken{ + accessToken: "valid-token", + expiresAt: now.Add(time.Second), + }, + } + cache.entries["idle-empty"] = &gigaChatTokenCacheEntry{} + cache.entries["in-flight-empty"] = &gigaChatTokenCacheEntry{refCount: 1} + + failedEntry := cache.acquireEntry("failed") + cache.releaseEntry("failed", failedEntry) + + entry := cache.acquireEntry("new") + entry.mu.Lock() + entry.token = gigaChatCachedToken{accessToken: "new-token", expiresAt: now.Add(time.Hour)} + entry.mu.Unlock() + cache.releaseEntry("new", entry) + + cache.mu.Lock() + defer cache.mu.Unlock() + if _, ok := cache.entries["expired"]; ok { + t.Fatal("expired cache entry was not pruned") + } + if _, ok := cache.entries["valid"]; !ok { + t.Fatal("valid cache entry was pruned") + } + if _, ok := cache.entries["idle-empty"]; ok { + t.Fatal("idle empty cache entry was not pruned") + } + if _, ok := cache.entries["in-flight-empty"]; !ok { + t.Fatal("in-flight empty cache entry was pruned") + } + if _, ok := cache.entries["failed"]; ok { + t.Fatal("failed token exchange cache entry was not released") + } + if _, ok := cache.entries["new"]; !ok { + t.Fatal("requested cache entry was not created") + } +} + +func testGigaChatTokenCacheConcurrentAccess(t *testing.T) { + t.Parallel() + + now := time.Unix(1_700_000_000, 0) + cache := newGigaChatTokenCache(func() time.Time { return now }) + + var wg sync.WaitGroup + for worker := 0; worker < 20; worker++ { + worker := worker + wg.Add(1) + go func() { + defer wg.Done() + for iteration := 0; iteration < 50; iteration++ { + cacheKey := "key-" + strconv.Itoa((worker+iteration)%7) + entry := cache.acquireEntry(cacheKey) + entry.mu.Lock() + entry.token = gigaChatCachedToken{ + accessToken: "token-" + strconv.Itoa(worker), + expiresAt: now.Add(time.Hour), + } + entry.mu.Unlock() + cache.releaseEntry(cacheKey, entry) + } + }() + } + wg.Wait() + + finalEntry := cache.acquireEntry("final") + cache.releaseEntry("final", finalEntry) + cache.mu.Lock() + defer cache.mu.Unlock() + for cacheKey, entry := range cache.entries { + entry.mu.Lock() + valid := entry.token.isValid(now) + entry.mu.Unlock() + if !valid { + t.Fatalf("unexpected stale cache entry after concurrent access: %s", cacheKey) + } + } +} + +func testGigaChatTokenCacheKeysDoNotReadCABundleFiles(t *testing.T) { + t.Parallel() + + certPEM, _ := generateGigaChatTestCertificate(t) + caBundleFile := writeGigaChatTestFile(t, "ca.pem", certPEM) + keyConfig := &schemas.GigaChatKeyConfig{CABundleFile: caBundleFile} + + oauthConfig := gigaChatOAuthConfig{ + authURL: "https://auth.example/token", + credentials: "test-credentials", + scope: schemas.DefaultGigaChatScope, + keyConfig: keyConfig, + } + oauthCacheKey := buildGigaChatOAuthCacheKey(oauthConfig) + + passwordConfig := gigaChatPasswordAuthConfig{ + tokenURL: "https://api.example/token", + user: "test-user", + password: "test-password", + keyConfig: keyConfig, + } + passwordCacheKey := buildGigaChatPasswordAuthCacheKey(passwordConfig) + + if err := os.Remove(caBundleFile); err != nil { + t.Fatalf("failed to remove CA bundle file: %v", err) + } + + repeatedOAuthCacheKey := buildGigaChatOAuthCacheKey(oauthConfig) + if repeatedOAuthCacheKey != oauthCacheKey { + t.Fatalf("OAuth token cache key changed after CA bundle removal: got %q, want %q", repeatedOAuthCacheKey, oauthCacheKey) + } + + repeatedPasswordCacheKey := buildGigaChatPasswordAuthCacheKey(passwordConfig) + if repeatedPasswordCacheKey != passwordCacheKey { + t.Fatalf("password token cache key changed after CA bundle removal: got %q, want %q", repeatedPasswordCacheKey, passwordCacheKey) + } +} + +func TestGigaChatAuthHeaders(t *testing.T) { + t.Parallel() + + t.Run("ExplicitAccessToken", testGigaChatAuthHeadersExplicitAccessToken) + t.Run("UserAgentLiteral", testGigaChatAuthHeadersUserAgentLiteral) + t.Run("KeyValueAccessToken", testGigaChatAuthHeadersKeyValueAccessToken) + t.Run("TLSOnlyOmitsBearerAuth", testGigaChatAuthHeadersTLSOnlyOmitsBearerAuth) + t.Run("CABundleOnlyDoesNotAuthenticate", testGigaChatAuthHeadersCABundleOnlyDoesNotAuthenticate) + t.Run("OAuthToken", testGigaChatAuthHeadersOAuthToken) + t.Run("BlocksProviderAuthorizationExtraHeader", testGigaChatAuthHeadersBlocksProviderAuthorizationExtraHeader) + t.Run("RejectsRequestAuthorizationExtraHeader", testGigaChatAuthHeadersRejectsRequestAuthorizationExtraHeader) + t.Run("PassesContextVars", testGigaChatAuthHeadersPassesContextVars) + t.Run("ForcedRefreshBypassesCachedOAuthToken", testGigaChatAuthHeadersForcedRefreshBypassesCachedOAuthToken) + t.Run("ForcedRefreshFallsBackFromExplicitTokenToOAuth", testGigaChatAuthHeadersForcedRefreshFallsBackFromExplicitTokenToOAuth) +} + +func testGigaChatAuthHeadersExplicitAccessToken(t *testing.T) { + t.Parallel() + + provider := newTestGigaChatProvider(t, time.Now) + headers, bifrostErr := provider.buildAuthHeaders(testBifrostContext(), schemas.Key{ + GigaChatKeyConfig: &schemas.GigaChatKeyConfig{ + AccessToken: schemas.NewSecretVar("explicit-access-token"), + }, + }) + if bifrostErr != nil { + t.Fatalf("buildAuthHeaders returned error: %v", bifrostErr) + } + assertGigaChatDefaultHeaders(t, headers, "Bearer explicit-access-token") +} + +func testGigaChatAuthHeadersUserAgentLiteral(t *testing.T) { + t.Parallel() + + provider := newTestGigaChatProvider(t, time.Now) + headers, bifrostErr := provider.buildAuthHeaders(testBifrostContext(), schemas.Key{ + GigaChatKeyConfig: &schemas.GigaChatKeyConfig{ + AccessToken: schemas.NewSecretVar("explicit-access-token"), + }, + }) + if bifrostErr != nil { + t.Fatalf("buildAuthHeaders returned error: %v", bifrostErr) + } + if got := headers[gigaChatUserAgentHeader]; got != "GigaChat-Bifrost-Provider" { + t.Fatalf("user-agent header mismatch: got %q, want %q", got, "GigaChat-Bifrost-Provider") + } +} + +func testGigaChatAuthHeadersKeyValueAccessToken(t *testing.T) { + t.Parallel() + + provider := newTestGigaChatProvider(t, time.Now) + headers, bifrostErr := provider.buildAuthHeaders(testBifrostContext(), schemas.Key{ + Value: *schemas.NewSecretVar("key-value-access-token"), + }) + if bifrostErr != nil { + t.Fatalf("buildAuthHeaders returned error: %v", bifrostErr) + } + assertGigaChatDefaultHeaders(t, headers, "Bearer key-value-access-token") +} + +func testGigaChatAuthHeadersTLSOnlyOmitsBearerAuth(t *testing.T) { + t.Parallel() + + provider := newTestGigaChatProvider(t, time.Now) + headers, bifrostErr := provider.buildAuthHeaders(testBifrostContext(), schemas.Key{ + GigaChatKeyConfig: &schemas.GigaChatKeyConfig{ + CertFile: "/secure/client.pem", + KeyFile: "/secure/client.key", + CABundleFile: "/secure/ca.pem", + }, + }) + if bifrostErr != nil { + t.Fatalf("buildAuthHeaders returned error: %v", bifrostErr) + } + assertGigaChatDefaultHeadersWithoutAuthorization(t, headers) +} + +func testGigaChatAuthHeadersCABundleOnlyDoesNotAuthenticate(t *testing.T) { + t.Parallel() + + provider := newTestGigaChatProvider(t, time.Now) + _, bifrostErr := provider.buildAuthHeaders(testBifrostContext(), schemas.Key{ + GigaChatKeyConfig: &schemas.GigaChatKeyConfig{ + CABundleFile: "/secure/ca.pem", + }, + }) + if bifrostErr == nil { + t.Fatal("expected error, got nil") + } + if !strings.Contains(bifrostErr.GetErrorString(), "mTLS cert_file/key_file") { + t.Fatalf("unexpected error: %v", bifrostErr) + } +} + +func testGigaChatAuthHeadersOAuthToken(t *testing.T) { + t.Parallel() + + now := time.Unix(1_700_000_000, 0) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"access_token":"oauth-access-token","expires_at":` + formatUnix(now.Add(30*time.Minute)) + `}`)) + })) + defer server.Close() + + provider := newTestGigaChatProvider(t, func() time.Time { return now }) + headers, bifrostErr := provider.buildAuthHeaders(testBifrostContext(), testGigaChatOAuthKey(server.URL, "", "test-credentials")) + if bifrostErr != nil { + t.Fatalf("buildAuthHeaders returned error: %v", bifrostErr) + } + assertGigaChatDefaultHeaders(t, headers, "Bearer oauth-access-token") +} + +func testGigaChatAuthHeadersBlocksProviderAuthorizationExtraHeader(t *testing.T) { + t.Parallel() + + provider := newTestGigaChatProvider(t, time.Now) + provider.networkConfig.ExtraHeaders = map[string]string{ + "authorization": "Bearer provider-authorization-token", + } + + _, bifrostErr := provider.buildAuthHeaders(testBifrostContext(), schemas.Key{ + GigaChatKeyConfig: &schemas.GigaChatKeyConfig{ + AccessToken: schemas.NewSecretVar("explicit-access-token"), + }, + }) + if bifrostErr == nil { + t.Fatal("expected error, got nil") + } + if !strings.Contains(bifrostErr.GetErrorString(), "extra_headers") { + t.Fatalf("unexpected error: %v", bifrostErr) + } + if strings.Contains(bifrostErr.GetErrorString(), "request extra headers") { + t.Fatalf("unexpected request header bypass hint: %v", bifrostErr) + } + assertNoGigaChatSecretLeak(t, bifrostErr.String()) +} + +func testGigaChatAuthHeadersRejectsRequestAuthorizationExtraHeader(t *testing.T) { + t.Parallel() + + var requestCount atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + requestCount.Add(1) + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"access_token":"oauth-access-token","expires_at":1893456000}`)) + })) + defer server.Close() + + ctx := testBifrostContext() + ctx.SetValue(schemas.BifrostContextKeyExtraHeaders, map[string][]string{ + "Authorization": {"Bearer context-authorization-token"}, + }) + + provider := newTestGigaChatProvider(t, time.Now) + _, bifrostErr := provider.buildAuthHeaders(ctx, testGigaChatOAuthKey(server.URL, "", "test-credentials")) + if bifrostErr == nil { + t.Fatal("expected error, got nil") + } + if !strings.Contains(bifrostErr.GetErrorString(), "request extra headers cannot include Authorization") { + t.Fatalf("unexpected error: %v", bifrostErr) + } + if requestCount.Load() != 0 { + t.Fatalf("request count mismatch: got %d, want 0", requestCount.Load()) + } + assertNoGigaChatSecretLeak(t, bifrostErr.String()) +} + +func testGigaChatAuthHeadersPassesContextVars(t *testing.T) { + t.Parallel() + + ctx := testBifrostContext() + ctx.SetValue(schemas.BifrostContextKeyExtraHeaders, map[string][]string{ + "X-Session-ID": {"session-id"}, + "X-Request-ID": {"request-id"}, + "X-Service-ID": {"service-id"}, + "X-Operation-ID": {"operation-id"}, + "X-Client-ID": {"client-id"}, + "X-Trace-ID": {"trace-id"}, + "X-Agent-ID": {"agent-id"}, + "X-Ignored-ID": {"ignored-id"}, + }) + + provider := newTestGigaChatProvider(t, time.Now) + provider.networkConfig.ExtraHeaders = map[string]string{ + "X-Service-ID": "provider-service-id", + } + + headers, bifrostErr := provider.buildAuthHeaders(ctx, schemas.Key{ + GigaChatKeyConfig: &schemas.GigaChatKeyConfig{ + AccessToken: schemas.NewSecretVar("explicit-access-token"), + }, + }) + if bifrostErr != nil { + t.Fatalf("buildAuthHeaders returned error: %v", bifrostErr) + } + + assertGigaChatDefaultHeaders(t, headers, "Bearer explicit-access-token") + expectedHeaders := map[string]string{ + "X-Session-ID": "session-id", + "X-Request-ID": "request-id", + "X-Service-ID": "service-id", + "X-Operation-ID": "operation-id", + "X-Client-ID": "client-id", + "X-Trace-ID": "trace-id", + "X-Agent-ID": "agent-id", + } + for key, want := range expectedHeaders { + if got := headers[key]; got != want { + t.Fatalf("%s mismatch: got %q, want %q", key, got, want) + } + } + if _, ok := headers["X-Ignored-ID"]; ok { + t.Fatalf("unexpected ignored header: %v", headers) + } +} + +func testGigaChatAuthHeadersForcedRefreshBypassesCachedOAuthToken(t *testing.T) { + t.Parallel() + + now := time.Unix(1_700_000_000, 0) + var requestCount atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + count := requestCount.Add(1) + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"access_token":"oauth-refresh-token-` + formatInt32(count) + `","expires_at":` + formatUnix(now.Add(30*time.Minute)) + `}`)) + })) + defer server.Close() + + provider := newTestGigaChatProvider(t, func() time.Time { return now }) + key := testGigaChatOAuthKey(server.URL, "", "test-credentials") + + headers, bifrostErr := provider.buildAuthHeaders(testBifrostContext(), key) + if bifrostErr != nil { + t.Fatalf("buildAuthHeaders returned error: %v", bifrostErr) + } + assertGigaChatDefaultHeaders(t, headers, "Bearer oauth-refresh-token-1") + + headers, bifrostErr = provider.refreshAuthHeaders(testBifrostContext(), key) + if bifrostErr != nil { + t.Fatalf("refreshAuthHeaders returned error: %v", bifrostErr) + } + assertGigaChatDefaultHeaders(t, headers, "Bearer oauth-refresh-token-2") + if requestCount.Load() != 2 { + t.Fatalf("request count mismatch: got %d, want 2", requestCount.Load()) + } +} + +func testGigaChatAuthHeadersForcedRefreshFallsBackFromExplicitTokenToOAuth(t *testing.T) { + t.Parallel() + + now := time.Unix(1_700_000_000, 0) + var requestCount atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + requestCount.Add(1) + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"access_token":"oauth-refreshed-token","expires_at":` + formatUnix(now.Add(30*time.Minute)) + `}`)) + })) + defer server.Close() + + provider := newTestGigaChatProvider(t, func() time.Time { return now }) + key := schemas.Key{ + GigaChatKeyConfig: &schemas.GigaChatKeyConfig{ + AccessToken: schemas.NewSecretVar("explicit-access-token"), + Credentials: schemas.NewSecretVar("test-credentials"), + AuthURL: server.URL, + }, + } + + headers, bifrostErr := provider.buildAuthHeaders(testBifrostContext(), key) + if bifrostErr != nil { + t.Fatalf("buildAuthHeaders returned error: %v", bifrostErr) + } + assertGigaChatDefaultHeaders(t, headers, "Bearer explicit-access-token") + if requestCount.Load() != 0 { + t.Fatalf("request count mismatch before refresh: got %d, want 0", requestCount.Load()) + } + + headers, bifrostErr = provider.refreshAuthHeaders(testBifrostContext(), key) + if bifrostErr != nil { + t.Fatalf("refreshAuthHeaders returned error: %v", bifrostErr) + } + assertGigaChatDefaultHeaders(t, headers, "Bearer oauth-refreshed-token") + if requestCount.Load() != 1 { + t.Fatalf("request count mismatch after refresh: got %d, want 1", requestCount.Load()) + } +} + +func testGigaChatOAuthRequestShapeAndDefaultScope(t *testing.T) { + t.Parallel() + + now := time.Unix(1_700_000_000, 0) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost { + t.Errorf("method mismatch: got %s", r.Method) + } + if r.URL.Path != "/api/v2/oauth" { + t.Errorf("path mismatch: got %s", r.URL.Path) + } + if contentType := r.Header.Get("Content-Type"); !strings.Contains(contentType, "application/x-www-form-urlencoded") { + t.Errorf("content type mismatch: got %q", contentType) + } + if accept := r.Header.Get("Accept"); accept != "application/json" { + t.Errorf("accept mismatch: got %q", accept) + } + if userAgent := r.Header.Get("User-Agent"); userAgent != gigaChatUserAgent { + t.Errorf("user-agent mismatch: got %q", userAgent) + } + if auth := r.Header.Get("Authorization"); auth != "Basic test-credentials" { + t.Errorf("authorization mismatch: got %q", auth) + } + requestID := r.Header.Get("RqUID") + parsedRequestID, err := uuid.Parse(requestID) + if err != nil { + t.Errorf("RqUID is not a UUID: %q", requestID) + } else if parsedRequestID.Version() != 4 { + t.Errorf("RqUID version mismatch: got %d", parsedRequestID.Version()) + } + + body, err := io.ReadAll(r.Body) + if err != nil { + t.Fatalf("failed to read request body: %v", err) + } + values, err := url.ParseQuery(string(body)) + if err != nil { + t.Fatalf("failed to parse request body: %v", err) + } + if scope := values.Get("scope"); scope != schemas.DefaultGigaChatScope { + t.Errorf("scope mismatch: got %q, want %q", scope, schemas.DefaultGigaChatScope) + } + + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"access_token":"token-1","expires_at":` + formatUnix(now.Add(30*time.Minute)) + `}`)) + })) + defer server.Close() + + provider := newTestGigaChatProvider(t, func() time.Time { return now }) + token, bifrostErr := provider.getOAuthAccessToken(testBifrostContext(), testGigaChatOAuthKey(server.URL+"/api/v2/oauth", "", "test-credentials")) + if bifrostErr != nil { + t.Fatalf("getOAuthAccessToken returned error: %v", bifrostErr) + } + if token != "token-1" { + t.Fatalf("token mismatch: got %q", token) + } +} + +func testGigaChatPasswordRequestShape(t *testing.T) { + t.Parallel() + + now := time.Unix(1_700_000_000, 0) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost { + t.Errorf("method mismatch: got %s", r.Method) + } + if r.URL.Path != "/api/v1/token" { + t.Errorf("path mismatch: got %s", r.URL.Path) + } + if contentType := r.Header.Get("Content-Type"); !strings.Contains(contentType, "application/x-www-form-urlencoded") { + t.Errorf("content type mismatch: got %q", contentType) + } + if accept := r.Header.Get("Accept"); accept != "application/json" { + t.Errorf("accept mismatch: got %q", accept) + } + if userAgent := r.Header.Get("User-Agent"); userAgent != gigaChatUserAgent { + t.Errorf("user-agent mismatch: got %q", userAgent) + } + wantAuth := "Basic " + base64.StdEncoding.EncodeToString([]byte("test-user:test-password")) + if auth := r.Header.Get("Authorization"); auth != wantAuth { + t.Errorf("authorization mismatch: got %q", auth) + } + + body, err := io.ReadAll(r.Body) + if err != nil { + t.Fatalf("failed to read request body: %v", err) + } + if len(body) != 0 { + t.Errorf("body mismatch: got %q, want empty body", string(body)) + } + + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"tok":"password-token-1","exp":` + formatUnixMilli(now.Add(30*time.Minute)) + `}`)) + })) + defer server.Close() + + provider := newTestGigaChatProvider(t, func() time.Time { return now }) + token, bifrostErr := provider.getPasswordAccessToken(testBifrostContext(), testGigaChatPasswordKey(server.URL+"/api", "test-user", "test-password")) + if bifrostErr != nil { + t.Fatalf("getPasswordAccessToken returned error: %v", bifrostErr) + } + if token != "password-token-1" { + t.Fatalf("token mismatch: got %q", token) + } +} + +func testGigaChatOAuthIgnoresClientCertificate(t *testing.T) { + t.Parallel() + + now := time.Unix(1_700_000_000, 0) + server, caBundleFile, certFile, keyFile := newGigaChatClientCertRequestingServer(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/api/v2/oauth" { + t.Errorf("unexpected path: %s", r.URL.Path) + w.WriteHeader(http.StatusNotFound) + return + } + if r.TLS != nil && len(r.TLS.PeerCertificates) != 0 { + t.Error("OAuth token request should not include client certificate") + w.WriteHeader(http.StatusBadRequest) + return + } + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"access_token":"oauth-token","expires_at":` + formatUnix(now.Add(30*time.Minute)) + `}`)) + })) + + provider := newTestGigaChatProvider(t, func() time.Time { return now }) + key := testGigaChatOAuthKey(server.URL+"/api/v2/oauth", "", "test-credentials") + key.GigaChatKeyConfig.CABundleFile = caBundleFile + key.GigaChatKeyConfig.CertFile = certFile + key.GigaChatKeyConfig.KeyFile = keyFile + + token, bifrostErr := provider.getOAuthAccessToken(testBifrostContext(), key) + if bifrostErr != nil { + t.Fatalf("getOAuthAccessToken returned error: %v", bifrostErr) + } + if token != "oauth-token" { + t.Fatalf("token mismatch: got %q", token) + } +} + +func testGigaChatPasswordIgnoresClientCertificate(t *testing.T) { + t.Parallel() + + now := time.Unix(1_700_000_000, 0) + server, caBundleFile, certFile, keyFile := newGigaChatClientCertRequestingServer(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/api/v1/token" { + t.Errorf("unexpected path: %s", r.URL.Path) + w.WriteHeader(http.StatusNotFound) + return + } + if r.TLS != nil && len(r.TLS.PeerCertificates) != 0 { + t.Error("password token request should not include client certificate") + w.WriteHeader(http.StatusBadRequest) + return + } + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"tok":"password-token","exp":` + formatUnixMilli(now.Add(30*time.Minute)) + `}`)) + })) + + provider := newTestGigaChatProvider(t, func() time.Time { return now }) + key := testGigaChatPasswordKey(server.URL+"/api", "test-user", "test-password") + key.GigaChatKeyConfig.CABundleFile = caBundleFile + key.GigaChatKeyConfig.CertFile = certFile + key.GigaChatKeyConfig.KeyFile = keyFile + + token, bifrostErr := provider.getPasswordAccessToken(testBifrostContext(), key) + if bifrostErr != nil { + t.Fatalf("getPasswordAccessToken returned error: %v", bifrostErr) + } + if token != "password-token" { + t.Fatalf("token mismatch: got %q", token) + } +} + +func newGigaChatMTLSServer(t *testing.T, handler http.Handler) (*httptest.Server, string, string, string) { + t.Helper() + + clientCertPEM, clientKeyPEM := generateGigaChatTestCertificate(t) + clientCAPool := x509.NewCertPool() + if !clientCAPool.AppendCertsFromPEM(clientCertPEM) { + t.Fatal("failed to parse client CA certificate") + } + + server := httptest.NewUnstartedServer(handler) + server.TLS = &tls.Config{ + MinVersion: tls.VersionTLS12, + ClientAuth: tls.RequireAndVerifyClientCert, + ClientCAs: clientCAPool, + } + server.StartTLS() + t.Cleanup(server.Close) + + serverCertPEM := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: server.Certificate().Raw}) + caBundleFile := writeGigaChatTestFile(t, "token-server-ca.pem", serverCertPEM) + certFile := writeGigaChatTestFile(t, "token-client.pem", clientCertPEM) + keyFile := writeGigaChatTestFile(t, "token-client.key", clientKeyPEM) + return server, caBundleFile, certFile, keyFile +} + +func newGigaChatClientCertRequestingServer(t *testing.T, handler http.Handler) (*httptest.Server, string, string, string) { + t.Helper() + + clientCertPEM, clientKeyPEM := generateGigaChatTestCertificate(t) + server := httptest.NewUnstartedServer(handler) + server.TLS = &tls.Config{ + MinVersion: tls.VersionTLS12, + ClientAuth: tls.RequestClientCert, + } + server.StartTLS() + t.Cleanup(server.Close) + + serverCertPEM := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: server.Certificate().Raw}) + caBundleFile := writeGigaChatTestFile(t, "token-server-ca.pem", serverCertPEM) + certFile := writeGigaChatTestFile(t, "token-client.pem", clientCertPEM) + keyFile := writeGigaChatTestFile(t, "token-client.key", clientKeyPEM) + return server, caBundleFile, certFile, keyFile +} + +func testGigaChatOAuthParsesMillisecondsExpiresAt(t *testing.T) { + t.Parallel() + + now := time.Unix(1_700_000_000, 0) + var requestCount atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + count := requestCount.Add(1) + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"access_token":"token-` + formatInt32(count) + `","expires_at":` + formatUnixMilli(now.Add(30*time.Minute)) + `}`)) + })) + defer server.Close() + + provider := newTestGigaChatProvider(t, func() time.Time { return now }) + key := testGigaChatOAuthKey(server.URL, "", "test-credentials") + + firstToken, bifrostErr := provider.getOAuthAccessToken(testBifrostContext(), key) + if bifrostErr != nil { + t.Fatalf("first getOAuthAccessToken returned error: %v", bifrostErr) + } + now = now.Add(31 * time.Minute) + secondToken, bifrostErr := provider.getOAuthAccessToken(testBifrostContext(), key) + if bifrostErr != nil { + t.Fatalf("second getOAuthAccessToken returned error: %v", bifrostErr) + } + + if firstToken != "token-1" || secondToken != "token-2" { + t.Fatalf("token refresh mismatch: first=%q second=%q", firstToken, secondToken) + } + if requestCount.Load() != 2 { + t.Fatalf("request count mismatch: got %d, want 2", requestCount.Load()) + } +} + +func testGigaChatPasswordParsesSecondsExpiresAt(t *testing.T) { + t.Parallel() + + now := time.Unix(1_700_000_000, 0) + var requestCount atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + count := requestCount.Add(1) + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"tok":"password-token-` + formatInt32(count) + `","exp":` + formatUnix(now.Add(30*time.Minute)) + `}`)) + })) + defer server.Close() + + provider := newTestGigaChatProvider(t, func() time.Time { return now }) + key := testGigaChatPasswordKey(server.URL, "test-user", "test-password") + + firstToken, bifrostErr := provider.getPasswordAccessToken(testBifrostContext(), key) + if bifrostErr != nil { + t.Fatalf("first getPasswordAccessToken returned error: %v", bifrostErr) + } + secondToken, bifrostErr := provider.getPasswordAccessToken(testBifrostContext(), key) + if bifrostErr != nil { + t.Fatalf("second getPasswordAccessToken returned error: %v", bifrostErr) + } + + if firstToken != "password-token-1" || secondToken != "password-token-1" { + t.Fatalf("cached token mismatch: first=%q second=%q", firstToken, secondToken) + } + if requestCount.Load() != 1 { + t.Fatalf("request count mismatch: got %d, want 1", requestCount.Load()) + } +} + +func testGigaChatPasswordCachesTokenBeforeLeeway(t *testing.T) { + t.Parallel() + + now := time.Unix(1_700_000_000, 0) + var requestCount atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + count := requestCount.Add(1) + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"tok":"password-token-` + formatInt32(count) + `","exp":` + formatUnixMilli(now.Add(30*time.Minute)) + `}`)) + })) + defer server.Close() + + provider := newTestGigaChatProvider(t, func() time.Time { return now }) + key := testGigaChatPasswordKey(server.URL, "test-user", "test-password") + + firstToken, bifrostErr := provider.getPasswordAccessToken(testBifrostContext(), key) + if bifrostErr != nil { + t.Fatalf("first getPasswordAccessToken returned error: %v", bifrostErr) + } + secondToken, bifrostErr := provider.getPasswordAccessToken(testBifrostContext(), key) + if bifrostErr != nil { + t.Fatalf("second getPasswordAccessToken returned error: %v", bifrostErr) + } + + if firstToken != "password-token-1" || secondToken != "password-token-1" { + t.Fatalf("cached token mismatch: first=%q second=%q", firstToken, secondToken) + } + if requestCount.Load() != 1 { + t.Fatalf("request count mismatch: got %d, want 1", requestCount.Load()) + } +} + +func testGigaChatPasswordCacheIncludesCABundle(t *testing.T) { + t.Parallel() + + now := time.Unix(1_700_000_000, 0) + var requestCount atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + count := requestCount.Add(1) + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"tok":"password-token-` + formatInt32(count) + `","exp":` + formatUnixMilli(now.Add(30*time.Minute)) + `}`)) + })) + defer server.Close() + + caBundlePEM1, _ := generateGigaChatTestCertificate(t) + caBundlePEM2, _ := generateGigaChatTestCertificate(t) + caBundleFile1 := writeGigaChatTestFile(t, "ca-1.pem", caBundlePEM1) + caBundleFile2 := writeGigaChatTestFile(t, "ca-2.pem", caBundlePEM2) + + provider := newTestGigaChatProvider(t, func() time.Time { return now }) + key1 := testGigaChatPasswordKey(server.URL, "test-user", "test-password") + key1.GigaChatKeyConfig.CABundleFile = caBundleFile1 + key2 := testGigaChatPasswordKey(server.URL, "test-user", "test-password") + key2.GigaChatKeyConfig.CABundleFile = caBundleFile2 + + firstToken, bifrostErr := provider.getPasswordAccessToken(testBifrostContext(), key1) + if bifrostErr != nil { + t.Fatalf("first getPasswordAccessToken returned error: %v", bifrostErr) + } + secondToken, bifrostErr := provider.getPasswordAccessToken(testBifrostContext(), key2) + if bifrostErr != nil { + t.Fatalf("second getPasswordAccessToken returned error: %v", bifrostErr) + } + thirdToken, bifrostErr := provider.getPasswordAccessToken(testBifrostContext(), key1) + if bifrostErr != nil { + t.Fatalf("third getPasswordAccessToken returned error: %v", bifrostErr) + } + + if firstToken != "password-token-1" || secondToken != "password-token-2" || thirdToken != "password-token-1" { + t.Fatalf("cache partition mismatch: first=%q second=%q third=%q", firstToken, secondToken, thirdToken) + } + if requestCount.Load() != 2 { + t.Fatalf("request count mismatch: got %d, want 2", requestCount.Load()) + } +} + +func testGigaChatPasswordRefreshesTokenInsideLeeway(t *testing.T) { + t.Parallel() + + now := time.Unix(1_700_000_000, 0) + var requestCount atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + count := requestCount.Add(1) + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"tok":"password-token-` + formatInt32(count) + `","exp":` + formatUnixMilli(now.Add(30*time.Minute)) + `}`)) + })) + defer server.Close() + + provider := newTestGigaChatProvider(t, func() time.Time { return now }) + key := testGigaChatPasswordKey(server.URL, "test-user", "test-password") + + firstToken, bifrostErr := provider.getPasswordAccessToken(testBifrostContext(), key) + if bifrostErr != nil { + t.Fatalf("first getPasswordAccessToken returned error: %v", bifrostErr) + } + now = now.Add(29*time.Minute + time.Second) + secondToken, bifrostErr := provider.getPasswordAccessToken(testBifrostContext(), key) + if bifrostErr != nil { + t.Fatalf("second getPasswordAccessToken returned error: %v", bifrostErr) + } + + if firstToken != "password-token-1" || secondToken != "password-token-2" { + t.Fatalf("token refresh mismatch: first=%q second=%q", firstToken, secondToken) + } + if requestCount.Load() != 2 { + t.Fatalf("request count mismatch: got %d, want 2", requestCount.Load()) + } +} + +func testGigaChatPasswordRejectsExpiredToken(t *testing.T) { + t.Parallel() + + now := time.Unix(1_700_000_000, 0) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"tok":"password-token-1","exp":` + formatUnixMilli(now.Add(-time.Second)) + `}`)) + })) + defer server.Close() + + provider := newTestGigaChatProvider(t, func() time.Time { return now }) + _, bifrostErr := provider.getPasswordAccessToken(testBifrostContext(), testGigaChatPasswordKey(server.URL, "super-secret-user", "super-secret-password")) + if bifrostErr == nil { + t.Fatal("expected error, got nil") + } + if !strings.Contains(bifrostErr.GetErrorString(), "already expired") { + t.Fatalf("unexpected error: %v", bifrostErr) + } + assertNoGigaChatSecretLeak(t, bifrostErr.String()) +} + +func testGigaChatPasswordHandlesProviderErrors(t *testing.T) { + t.Parallel() + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusUnauthorized) + _, _ = w.Write([]byte(`{"code":4,"message":"invalid password auth"}`)) + })) + defer server.Close() + + provider := newTestGigaChatProvider(t, time.Now) + _, bifrostErr := provider.getPasswordAccessToken(testBifrostContext(), testGigaChatPasswordKey(server.URL, "super-secret-user", "super-secret-password")) + if bifrostErr == nil { + t.Fatal("expected error, got nil") + } + if bifrostErr.Error == nil || bifrostErr.Error.Message != "invalid password auth" { + t.Fatalf("unexpected error: %v", bifrostErr) + } + assertGigaChatTokenCacheEmpty(t, provider.tokenCache) + assertNoGigaChatSecretLeak(t, bifrostErr.String()) +} + +func testGigaChatPasswordHandlesMalformedResponses(t *testing.T) { + t.Parallel() + + testCases := []struct { + name string + body string + }{ + {name: "invalid json", body: `not-json`}, + {name: "missing token", body: `{"exp":1893456000000}`}, + {name: "missing expiry", body: `{"tok":"password-token-1"}`}, + } + + for _, testCase := range testCases { + testCase := testCase + t.Run(testCase.name, func(t *testing.T) { + t.Parallel() + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(testCase.body)) + })) + defer server.Close() + + provider := newTestGigaChatProvider(t, func() time.Time { return time.Unix(100, 0) }) + _, bifrostErr := provider.getPasswordAccessToken(testBifrostContext(), testGigaChatPasswordKey(server.URL, "super-secret-user", "super-secret-password")) + if bifrostErr == nil { + t.Fatal("expected error, got nil") + } + assertNoGigaChatSecretLeak(t, bifrostErr.String()) + }) + } +} + +func testGigaChatPasswordMissingUserPassword(t *testing.T) { + t.Parallel() + + provider := newTestGigaChatProvider(t, time.Now) + testCases := []struct { + name string + key schemas.Key + }{ + { + name: "missing password", + key: schemas.Key{ + GigaChatKeyConfig: &schemas.GigaChatKeyConfig{ + User: schemas.NewSecretVar("test-user"), + }, + }, + }, + { + name: "missing user", + key: schemas.Key{ + GigaChatKeyConfig: &schemas.GigaChatKeyConfig{ + Password: schemas.NewSecretVar("test-password"), + }, + }, + }, + { + name: "empty user env", + key: schemas.Key{ + GigaChatKeyConfig: &schemas.GigaChatKeyConfig{ + User: schemas.NewSecretVar("env.MISSING_GIGACHAT_USER_FOR_TEST"), + Password: schemas.NewSecretVar("test-password"), + }, + }, + }, + { + name: "empty password env", + key: schemas.Key{ + GigaChatKeyConfig: &schemas.GigaChatKeyConfig{ + User: schemas.NewSecretVar("test-user"), + Password: schemas.NewSecretVar("env.MISSING_GIGACHAT_PASSWORD_FOR_TEST"), + }, + }, + }, + } + + for _, testCase := range testCases { + testCase := testCase + t.Run(testCase.name, func(t *testing.T) { + t.Parallel() + + _, bifrostErr := provider.getPasswordAccessToken(testBifrostContext(), testCase.key) + if bifrostErr == nil { + t.Fatal("expected error, got nil") + } + assertNoGigaChatSecretLeak(t, bifrostErr.String()) + }) + } +} + +func testGigaChatAuthPriority(t *testing.T) { + t.Parallel() + + t.Run("AccessTokenBeforeTokenFlows", func(t *testing.T) { + t.Parallel() + + provider := newTestGigaChatProvider(t, time.Now) + token, bifrostErr := provider.getGigaChatAccessToken(testBifrostContext(), schemas.Key{ + GigaChatKeyConfig: &schemas.GigaChatKeyConfig{ + AccessToken: schemas.NewSecretVar("explicit-token"), + Credentials: schemas.NewSecretVar("test-credentials"), + User: schemas.NewSecretVar("test-user"), + Password: schemas.NewSecretVar("test-password"), + }, + }) + if bifrostErr != nil { + t.Fatalf("getGigaChatAccessToken returned error: %v", bifrostErr) + } + if token != "explicit-token" { + t.Fatalf("token mismatch: got %q", token) + } + }) + + t.Run("OAuthBeforePassword", func(t *testing.T) { + t.Parallel() + + now := time.Unix(1_700_000_000, 0) + var oauthRequests atomic.Int32 + var passwordRequests atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + switch r.URL.Path { + case "/api/v2/oauth": + oauthRequests.Add(1) + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"access_token":"oauth-token","expires_at":` + formatUnix(now.Add(30*time.Minute)) + `}`)) + case "/api/v1/token": + passwordRequests.Add(1) + w.WriteHeader(http.StatusInternalServerError) + default: + t.Errorf("unexpected path: %s", r.URL.Path) + w.WriteHeader(http.StatusNotFound) + } + })) + defer server.Close() + + provider := newTestGigaChatProvider(t, func() time.Time { return now }) + token, bifrostErr := provider.getGigaChatAccessToken(testBifrostContext(), schemas.Key{ + GigaChatKeyConfig: &schemas.GigaChatKeyConfig{ + Credentials: schemas.NewSecretVar("test-credentials"), + AuthURL: server.URL + "/api/v2/oauth", + BaseURL: server.URL + "/api", + User: schemas.NewSecretVar("test-user"), + Password: schemas.NewSecretVar("test-password"), + }, + }) + if bifrostErr != nil { + t.Fatalf("getGigaChatAccessToken returned error: %v", bifrostErr) + } + if token != "oauth-token" { + t.Fatalf("token mismatch: got %q", token) + } + if oauthRequests.Load() != 1 { + t.Fatalf("oauth request count mismatch: got %d, want 1", oauthRequests.Load()) + } + if passwordRequests.Load() != 0 { + t.Fatalf("password request count mismatch: got %d, want 0", passwordRequests.Load()) + } + }) +} + +func testGigaChatOAuthCachesTokenBeforeLeeway(t *testing.T) { + t.Parallel() + + now := time.Unix(1_700_000_000, 0) + var requestCount atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + count := requestCount.Add(1) + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"access_token":"token-` + formatInt32(count) + `","expires_at":` + formatUnix(now.Add(30*time.Minute)) + `}`)) + })) + defer server.Close() + + provider := newTestGigaChatProvider(t, func() time.Time { return now }) + key := testGigaChatOAuthKey(server.URL, "", "test-credentials") + + firstToken, bifrostErr := provider.getOAuthAccessToken(testBifrostContext(), key) + if bifrostErr != nil { + t.Fatalf("first getOAuthAccessToken returned error: %v", bifrostErr) + } + secondToken, bifrostErr := provider.getOAuthAccessToken(testBifrostContext(), key) + if bifrostErr != nil { + t.Fatalf("second getOAuthAccessToken returned error: %v", bifrostErr) + } + + if firstToken != "token-1" || secondToken != "token-1" { + t.Fatalf("cached token mismatch: first=%q second=%q", firstToken, secondToken) + } + if requestCount.Load() != 1 { + t.Fatalf("request count mismatch: got %d, want 1", requestCount.Load()) + } +} + +func testGigaChatOAuthCacheIncludesCABundle(t *testing.T) { + t.Parallel() + + now := time.Unix(1_700_000_000, 0) + var requestCount atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + count := requestCount.Add(1) + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"access_token":"token-` + formatInt32(count) + `","expires_at":` + formatUnix(now.Add(30*time.Minute)) + `}`)) + })) + defer server.Close() + + caBundlePEM1, _ := generateGigaChatTestCertificate(t) + caBundlePEM2, _ := generateGigaChatTestCertificate(t) + caBundleFile1 := writeGigaChatTestFile(t, "ca-1.pem", caBundlePEM1) + caBundleFile2 := writeGigaChatTestFile(t, "ca-2.pem", caBundlePEM2) + + provider := newTestGigaChatProvider(t, func() time.Time { return now }) + key1 := testGigaChatOAuthKey(server.URL, "", "test-credentials") + key1.GigaChatKeyConfig.CABundleFile = caBundleFile1 + key2 := testGigaChatOAuthKey(server.URL, "", "test-credentials") + key2.GigaChatKeyConfig.CABundleFile = caBundleFile2 + + firstToken, bifrostErr := provider.getOAuthAccessToken(testBifrostContext(), key1) + if bifrostErr != nil { + t.Fatalf("first getOAuthAccessToken returned error: %v", bifrostErr) + } + secondToken, bifrostErr := provider.getOAuthAccessToken(testBifrostContext(), key2) + if bifrostErr != nil { + t.Fatalf("second getOAuthAccessToken returned error: %v", bifrostErr) + } + thirdToken, bifrostErr := provider.getOAuthAccessToken(testBifrostContext(), key1) + if bifrostErr != nil { + t.Fatalf("third getOAuthAccessToken returned error: %v", bifrostErr) + } + + if firstToken != "token-1" || secondToken != "token-2" || thirdToken != "token-1" { + t.Fatalf("cache partition mismatch: first=%q second=%q third=%q", firstToken, secondToken, thirdToken) + } + if requestCount.Load() != 2 { + t.Fatalf("request count mismatch: got %d, want 2", requestCount.Load()) + } +} + +func testGigaChatOAuthRefreshesTokenInsideLeeway(t *testing.T) { + t.Parallel() + + now := time.Unix(1_700_000_000, 0) + var requestCount atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + count := requestCount.Add(1) + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"access_token":"token-` + formatInt32(count) + `","expires_at":` + formatUnix(now.Add(30*time.Minute)) + `}`)) + })) + defer server.Close() + + provider := newTestGigaChatProvider(t, func() time.Time { return now }) + key := testGigaChatOAuthKey(server.URL, "", "test-credentials") + + firstToken, bifrostErr := provider.getOAuthAccessToken(testBifrostContext(), key) + if bifrostErr != nil { + t.Fatalf("first getOAuthAccessToken returned error: %v", bifrostErr) + } + now = now.Add(29*time.Minute + time.Second) + secondToken, bifrostErr := provider.getOAuthAccessToken(testBifrostContext(), key) + if bifrostErr != nil { + t.Fatalf("second getOAuthAccessToken returned error: %v", bifrostErr) + } + + if firstToken != "token-1" || secondToken != "token-2" { + t.Fatalf("token refresh mismatch: first=%q second=%q", firstToken, secondToken) + } + if requestCount.Load() != 2 { + t.Fatalf("request count mismatch: got %d, want 2", requestCount.Load()) + } +} + +func testGigaChatOAuthHandlesProviderErrors(t *testing.T) { + t.Parallel() + + testCases := []struct { + name string + statusCode int + body string + wantMessage string + wantCode string + }{ + { + name: "bad request code shape", + statusCode: http.StatusBadRequest, + body: `{"code":5,"message":"scope is empty"}`, + wantMessage: "scope is empty", + wantCode: "5", + }, + { + name: "unauthorized code shape", + statusCode: http.StatusUnauthorized, + body: `{"code":4,"message":"Can't decode 'Authorization' header"}`, + wantMessage: "Can't decode 'Authorization' header", + wantCode: "4", + }, + { + name: "server status shape", + statusCode: http.StatusInternalServerError, + body: `{"status":500,"message":"Internal Server Error"}`, + wantMessage: "Internal Server Error", + wantCode: "500", + }, + } + + for _, testCase := range testCases { + testCase := testCase + t.Run(testCase.name, func(t *testing.T) { + t.Parallel() + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(testCase.statusCode) + _, _ = w.Write([]byte(testCase.body)) + })) + defer server.Close() + + provider := newTestGigaChatProvider(t, time.Now) + _, bifrostErr := provider.getOAuthAccessToken(testBifrostContext(), testGigaChatOAuthKey(server.URL, "", "super-secret-credentials")) + if bifrostErr == nil { + t.Fatal("expected error, got nil") + } + if bifrostErr.Error == nil { + t.Fatal("expected error field, got nil") + } + if bifrostErr.Error.Message != testCase.wantMessage { + t.Fatalf("message mismatch: got %q, want %q", bifrostErr.Error.Message, testCase.wantMessage) + } + if bifrostErr.Error.Code == nil || *bifrostErr.Error.Code != testCase.wantCode { + t.Fatalf("code mismatch: got %#v, want %q", bifrostErr.Error.Code, testCase.wantCode) + } + assertGigaChatTokenCacheEmpty(t, provider.tokenCache) + assertNoGigaChatSecretLeak(t, bifrostErr.String()) + }) + } +} + +func assertGigaChatTokenCacheEmpty(t *testing.T, cache *gigaChatTokenCache) { + t.Helper() + cache.mu.Lock() + defer cache.mu.Unlock() + if len(cache.entries) != 0 { + t.Fatalf("token cache retained %d entries after failed exchange", len(cache.entries)) + } +} + +func testGigaChatOAuthHandlesMalformedResponses(t *testing.T) { + t.Parallel() + + testCases := []struct { + name string + body string + }{ + {name: "invalid json", body: `not-json`}, + {name: "missing token", body: `{"expires_at":1893456000}`}, + {name: "missing expiry", body: `{"access_token":"token-1"}`}, + {name: "expired token", body: `{"access_token":"token-1","expires_at":1}`}, + } + + for _, testCase := range testCases { + testCase := testCase + t.Run(testCase.name, func(t *testing.T) { + t.Parallel() + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(testCase.body)) + })) + defer server.Close() + + provider := newTestGigaChatProvider(t, func() time.Time { return time.Unix(100, 0) }) + _, bifrostErr := provider.getOAuthAccessToken(testBifrostContext(), testGigaChatOAuthKey(server.URL, "", "super-secret-credentials")) + if bifrostErr == nil { + t.Fatal("expected error, got nil") + } + assertNoGigaChatSecretLeak(t, bifrostErr.String()) + }) + } +} + +func testGigaChatOAuthMissingCredentials(t *testing.T) { + t.Parallel() + + provider := newTestGigaChatProvider(t, time.Now) + _, bifrostErr := provider.getOAuthAccessToken(testBifrostContext(), schemas.Key{ + GigaChatKeyConfig: &schemas.GigaChatKeyConfig{}, + }) + if bifrostErr == nil { + t.Fatal("expected error, got nil") + } + if !strings.Contains(bifrostErr.GetErrorString(), "credentials") { + t.Fatalf("unexpected error: %v", bifrostErr) + } + + _, bifrostErr = provider.getOAuthAccessToken(testBifrostContext(), schemas.Key{ + GigaChatKeyConfig: &schemas.GigaChatKeyConfig{ + Credentials: schemas.NewSecretVar("env.MISSING_GIGACHAT_CREDENTIALS_FOR_TEST"), + }, + }) + if bifrostErr == nil { + t.Fatal("expected unresolved env error, got nil") + } + if !strings.Contains(bifrostErr.GetErrorString(), "empty value") { + t.Fatalf("unexpected unresolved env error: %v", bifrostErr) + } +} + +func testGigaChatOAuthContextCancellation(t *testing.T) { + t.Parallel() + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + time.Sleep(10 * time.Millisecond) + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"access_token":"token-1","expires_at":1893456000}`)) + })) + defer server.Close() + + cancelledContext, cancel := context.WithCancel(context.Background()) + cancel() + + provider := newTestGigaChatProvider(t, time.Now) + _, bifrostErr := provider.getOAuthAccessToken(schemas.NewBifrostContext(cancelledContext, schemas.NoDeadline), testGigaChatOAuthKey(server.URL, "", "test-credentials")) + if bifrostErr == nil { + t.Fatal("expected cancellation error, got nil") + } + if bifrostErr.Error == nil || bifrostErr.Error.Type == nil || *bifrostErr.Error.Type != schemas.RequestCancelled { + t.Fatalf("unexpected cancellation error: %v", bifrostErr) + } +} + +func newTestGigaChatProvider(t *testing.T, now func() time.Time) *GigaChatProvider { + t.Helper() + + provider, err := NewGigaChatProvider(&schemas.ProviderConfig{}, nil) + if err != nil { + t.Fatalf("NewGigaChatProvider returned error: %v", err) + } + dialer := &net.Dialer{} + provider.client.Dial = func(addr string) (net.Conn, error) { + return dialer.Dial("tcp", addr) + } + provider.client.DialTimeout = nil + provider.streamingClient.Dial = provider.client.Dial + provider.streamingClient.DialTimeout = nil + provider.tokenCache = newGigaChatTokenCache(now) + return provider +} + +func testGigaChatOAuthKey(authURL string, scope string, credentials string) schemas.Key { + return schemas.Key{ + GigaChatKeyConfig: &schemas.GigaChatKeyConfig{ + Credentials: schemas.NewSecretVar(credentials), + Scope: scope, + AuthURL: authURL, + }, + } +} + +func testGigaChatPasswordKey(baseURL string, user string, password string) schemas.Key { + return schemas.Key{ + GigaChatKeyConfig: &schemas.GigaChatKeyConfig{ + BaseURL: baseURL, + User: schemas.NewSecretVar(user), + Password: schemas.NewSecretVar(password), + }, + } +} + +func testBifrostContext() *schemas.BifrostContext { + return schemas.NewBifrostContext(context.Background(), schemas.NoDeadline) +} + +func assertNoGigaChatSecretLeak(t *testing.T, output string) { + t.Helper() + for _, secret := range []string{"super-secret-credentials", "test-credentials", "super-secret-user", "super-secret-password", "test-user", "test-password", "explicit-access-token", "key-value-access-token", "provider-authorization-token", "context-authorization-token"} { + if strings.Contains(output, secret) { + t.Fatalf("secret %q leaked in %s", secret, output) + } + } +} + +func assertGigaChatDefaultHeaders(t *testing.T, headers map[string]string, wantAuthorization string) { + t.Helper() + + if got := headers[gigaChatAuthorizationHeader]; got != wantAuthorization { + t.Fatalf("authorization header mismatch: got %q, want %q", got, wantAuthorization) + } + if got := headers[gigaChatUserAgentHeader]; got != gigaChatUserAgent { + t.Fatalf("user-agent header mismatch: got %q, want %q", got, gigaChatUserAgent) + } +} + +func assertGigaChatDefaultHeadersWithoutAuthorization(t *testing.T, headers map[string]string) { + t.Helper() + + if got := headers[gigaChatAuthorizationHeader]; got != "" { + t.Fatalf("unexpected authorization header: %q", got) + } + if got := headers[gigaChatUserAgentHeader]; got != gigaChatUserAgent { + t.Fatalf("user-agent header mismatch: got %q, want %q", got, gigaChatUserAgent) + } +} + +func formatUnix(value time.Time) string { + return formatInt64(value.Unix()) +} + +func formatUnixMilli(value time.Time) string { + return formatInt64(value.UnixMilli()) +} + +func formatInt32(value int32) string { + return formatInt64(int64(value)) +} + +func formatInt64(value int64) string { + return strconv.FormatInt(value, 10) +} diff --git a/core/providers/gigachat/batch_test.go b/core/providers/gigachat/batch_test.go new file mode 100644 index 00000000000..6e503aee58e --- /dev/null +++ b/core/providers/gigachat/batch_test.go @@ -0,0 +1,951 @@ +package gigachat + +import ( + "bytes" + "encoding/json" + "io" + "net/http" + "net/http/httptest" + "strings" + "sync/atomic" + "testing" + + schemas "github.com/maximhq/bifrost/core/schemas" +) + +func TestGigaChatBatchTypesJSON(t *testing.T) { + t.Parallel() + + resultFileID := "file-result" + raw, err := json.Marshal(GigaChatBatches{Data: []GigaChatBatch{{ + ID: "batch-1", + Object: "batch", + Method: GigaChatBatchMethodChatCompletions, + Status: GigaChatBatchStatusCompleted, + ResultFileID: &resultFileID, + RequestCounts: &GigaChatBatchRequestCounts{ + Total: 3, + Completed: 2, + Failed: 1, + }, + }}}) + if err != nil { + t.Fatalf("marshal GigaChatBatches returned error: %v", err) + } + + var decoded GigaChatBatches + if err := json.Unmarshal(raw, &decoded); err != nil { + t.Fatalf("unmarshal GigaChatBatches returned error: %v", err) + } + if len(decoded.Data) != 1 { + t.Fatalf("decoded %d batches, want 1", len(decoded.Data)) + } + batch := decoded.Data[0] + if batch.ID != "batch-1" || batch.Method != GigaChatBatchMethodChatCompletions || batch.Status != GigaChatBatchStatusCompleted { + t.Fatalf("decoded batch mismatch: %#v", batch) + } + if batch.RequestCounts == nil || batch.RequestCounts.Total != 3 || batch.RequestCounts.Completed != 2 || batch.RequestCounts.Failed != 1 { + t.Fatalf("request counts mismatch: %#v", batch.RequestCounts) + } +} + +func TestGigaChatBatchesUnmarshalJSON(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + payload string + wantLen int + wantID string + }{ + { + name: "data wrapper", + payload: `{"data":[{"id":"batch-wrapper","object":"batch","method":"chat_completions","status":"created"}]}`, + wantLen: 1, + wantID: "batch-wrapper", + }, + { + name: "root array", + payload: `[{"id":"batch-array","object":"batch","method":"embedder","status":"completed"}]`, + wantLen: 1, + wantID: "batch-array", + }, + { + name: "empty root array", + payload: `[]`, + wantLen: 0, + }, + } + + for _, tt := range tests { + tt := tt + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + var decoded GigaChatBatches + if err := json.Unmarshal([]byte(tt.payload), &decoded); err != nil { + t.Fatalf("unmarshal GigaChatBatches returned error: %v", err) + } + if len(decoded.Data) != tt.wantLen { + t.Fatalf("decoded %d batches, want %d", len(decoded.Data), tt.wantLen) + } + if tt.wantLen > 0 && decoded.Data[0].ID != tt.wantID { + t.Fatalf("decoded id %q, want %q", decoded.Data[0].ID, tt.wantID) + } + }) + } +} + +func TestToGigaChatBatchMethod(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + endpoint schemas.BatchEndpoint + want GigaChatBatchMethod + wantErr string + }{ + {name: "chat completions", endpoint: schemas.BatchEndpointChatCompletions, want: GigaChatBatchMethodChatCompletions}, + {name: "chat completions without version", endpoint: "/chat/completions", want: GigaChatBatchMethodChatCompletions}, + {name: "responses", endpoint: schemas.BatchEndpointResponses, want: GigaChatBatchMethodResponses}, + {name: "responses without version", endpoint: "/responses", want: GigaChatBatchMethodResponses}, + {name: "embeddings", endpoint: schemas.BatchEndpointEmbeddings, want: GigaChatBatchMethodEmbedder}, + {name: "embeddings without version", endpoint: "/embeddings", want: GigaChatBatchMethodEmbedder}, + {name: "unknown", endpoint: schemas.BatchEndpointCompletions, wantErr: "do not support endpoint"}, + } + + for _, tt := range tests { + tt := tt + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + got, err := toGigaChatBatchMethod(tt.endpoint) + if tt.wantErr != "" { + if err == nil || !strings.Contains(err.Error(), tt.wantErr) { + t.Fatalf("toGigaChatBatchMethod() error = %v, want containing %q", err, tt.wantErr) + } + return + } + if err != nil { + t.Fatalf("toGigaChatBatchMethod() returned error: %v", err) + } + if got != tt.want { + t.Fatalf("toGigaChatBatchMethod() = %q, want %q", got, tt.want) + } + }) + } +} + +func TestToBifrostGigaChatBatchEndpoint(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + method GigaChatBatchMethod + want string + }{ + {name: "chat completions", method: GigaChatBatchMethodChatCompletions, want: string(schemas.BatchEndpointChatCompletions)}, + {name: "responses", method: GigaChatBatchMethodResponses, want: string(schemas.BatchEndpointResponses)}, + {name: "embeddings", method: GigaChatBatchMethodEmbedder, want: string(schemas.BatchEndpointEmbeddings)}, + {name: "unknown", method: GigaChatBatchMethod("unknown"), want: ""}, + } + + for _, tt := range tests { + tt := tt + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + if got := toBifrostGigaChatBatchEndpoint(tt.method); got != tt.want { + t.Fatalf("toBifrostGigaChatBatchEndpoint(%q) = %q, want %q", tt.method, got, tt.want) + } + }) + } +} + +func TestToBifrostGigaChatBatchStatus(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + status GigaChatBatchStatus + want schemas.BatchStatus + }{ + {name: "created maps to validating", status: GigaChatBatchStatusCreated, want: schemas.BatchStatusValidating}, + {name: "in progress", status: GigaChatBatchStatusInProgress, want: schemas.BatchStatusInProgress}, + {name: "completed", status: GigaChatBatchStatusCompleted, want: schemas.BatchStatusCompleted}, + {name: "unknown preserved", status: GigaChatBatchStatus("queued"), want: schemas.BatchStatus("queued")}, + } + + for _, tt := range tests { + tt := tt + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + if got := toBifrostGigaChatBatchStatus(tt.status); got != tt.want { + t.Fatalf("toBifrostGigaChatBatchStatus(%q) = %q, want %q", tt.status, got, tt.want) + } + }) + } +} + +func TestConvertGigaChatBatchInputJSONL(t *testing.T) { + t.Parallel() + + t.Run("ChatCompletions", testConvertGigaChatBatchInputJSONLChatCompletions) + t.Run("Responses", testConvertGigaChatBatchInputJSONLResponses) + t.Run("Embeddings", testConvertGigaChatBatchInputJSONLEmbeddings) + t.Run("InlineRequestItems", testConvertGigaChatBatchRequestItemsToJSONL) + t.Run("UnsupportedEndpoint", testConvertGigaChatBatchInputJSONLUnsupportedEndpoint) +} + +func TestGigaChatBatchesHTTP(t *testing.T) { + t.Parallel() + + t.Run("CreateTransformsFileRows", testGigaChatBatchCreateTransformsFileRows) + t.Run("CreateUsesKeyBaseURLAndRefreshesTokenAfterUnauthorized", testGigaChatBatchCreateUsesKeyBaseURLAndRefreshesTokenAfterUnauthorized) + t.Run("ListParsesWrapper", testGigaChatBatchListParsesWrapper) + t.Run("ListPaginatesLocally", testGigaChatBatchListPaginatesLocally) + t.Run("ListParsesEmptyRootArray", testGigaChatBatchListParsesEmptyRootArray) + t.Run("RetrieveParsesSingleObject", testGigaChatBatchRetrieveParsesSingleObject) + t.Run("RetrieveMapsResultFileID", testGigaChatBatchRetrieveMapsResultFileID) + t.Run("PreservesResponsesEndpoint", testGigaChatBatchPreservesResponsesEndpoint) + t.Run("ResultsDownloadsOutputFile", testGigaChatBatchResultsDownloadsOutputFile) + t.Run("ResultsWithoutOutputFileID", testGigaChatBatchResultsWithoutOutputFileID) + t.Run("OutputFileRejectsEmptyKeys", testGigaChatBatchOutputFileRejectsEmptyKeys) + t.Run("UnsupportedEndpoint", testGigaChatBatchCreateUnsupportedEndpoint) + t.Run("UnsupportedCompletionWindow", testGigaChatBatchCreateUnsupportedCompletionWindow) +} + +func testConvertGigaChatBatchInputJSONLChatCompletions(t *testing.T) { + t.Parallel() + + input := []byte(`{"custom_id":"chat-1","method":"POST","url":"/v1/chat/completions","body":{"model":"GigaChat","messages":[{"role":"user","content":"Hello"}],"max_tokens":64,"temperature":0.2}}` + "\n") + output, err := convertGigaChatBatchInputJSONL(schemas.BatchEndpointChatCompletions, input) + if err != nil { + t.Fatalf("convertGigaChatBatchInputJSONL returned error: %v", err) + } + + row := decodeGigaChatBatchTestRow(t, output) + if row.ID != "chat-1" { + t.Fatalf("row id mismatch: got %q", row.ID) + } + var request GigaChatChatRequest + if err := json.Unmarshal(row.Request, &request); err != nil { + t.Fatalf("unmarshal GigaChat chat request: %v", err) + } + if request.Model != "GigaChat" || len(request.Messages) != 1 || request.Messages[0].Role != "user" { + t.Fatalf("chat request mismatch: %#v", request) + } + if request.Messages[0].Content == nil || request.Messages[0].Content.ContentStr == nil || *request.Messages[0].Content.ContentStr != "Hello" { + t.Fatalf("chat content mismatch: %#v", request.Messages[0].Content) + } + if request.MaxTokens == nil || *request.MaxTokens != 64 { + t.Fatalf("max_tokens mismatch: %#v", request.MaxTokens) + } + if request.Temperature == nil || *request.Temperature != 0.2 { + t.Fatalf("temperature mismatch: %#v", request.Temperature) + } + if request.Stream == nil || *request.Stream { + t.Fatalf("stream mismatch: %#v", request.Stream) + } +} + +func testConvertGigaChatBatchInputJSONLResponses(t *testing.T) { + t.Parallel() + + input := []byte(`{"custom_id":"resp-1","method":"POST","url":"/v1/responses","body":{"model":"GigaChat-2","input":"Summarize this","instructions":"Be concise.","max_output_tokens":32}}` + "\n") + output, err := convertGigaChatBatchInputJSONL(schemas.BatchEndpointResponses, input) + if err != nil { + t.Fatalf("convertGigaChatBatchInputJSONL returned error: %v", err) + } + + row := decodeGigaChatBatchTestRow(t, output) + if row.ID != "resp-1" { + t.Fatalf("row id mismatch: got %q", row.ID) + } + var request GigaChatResponsesRequest + if err := json.Unmarshal(row.Request, &request); err != nil { + t.Fatalf("unmarshal GigaChat responses request: %v", err) + } + if request.Model != "GigaChat-2" { + t.Fatalf("model mismatch: got %q", request.Model) + } + if len(request.Messages) != 2 { + t.Fatalf("message count mismatch: got %d", len(request.Messages)) + } + if request.Messages[0].Role != "system" || request.Messages[0].Content[0].Text == nil || *request.Messages[0].Content[0].Text != "Be concise." { + t.Fatalf("instruction message mismatch: %#v", request.Messages[0]) + } + if request.Messages[1].Role != "user" || request.Messages[1].Content[0].Text == nil || *request.Messages[1].Content[0].Text != "Summarize this" { + t.Fatalf("input message mismatch: %#v", request.Messages[1]) + } + if request.ModelOptions == nil || request.ModelOptions.MaxTokens == nil || *request.ModelOptions.MaxTokens != 32 { + t.Fatalf("max output tokens mismatch: %#v", request.ModelOptions) + } +} + +func testConvertGigaChatBatchInputJSONLEmbeddings(t *testing.T) { + t.Parallel() + + input := []byte(`{"custom_id":"emb-1","method":"POST","url":"/v1/embeddings","body":{"model":"Embeddings","input":["first","second"]}}` + "\n") + output, err := convertGigaChatBatchInputJSONL(schemas.BatchEndpointEmbeddings, input) + if err != nil { + t.Fatalf("convertGigaChatBatchInputJSONL returned error: %v", err) + } + + row := decodeGigaChatBatchTestRow(t, output) + if row.ID != "emb-1" { + t.Fatalf("row id mismatch: got %q", row.ID) + } + var request GigaChatEmbeddingRequest + if err := json.Unmarshal(row.Request, &request); err != nil { + t.Fatalf("unmarshal GigaChat embedding request: %v", err) + } + if request.Model != "Embeddings" { + t.Fatalf("model mismatch: got %q", request.Model) + } + if request.Input == nil || len(request.Input.Texts) != 2 || request.Input.Texts[0] != "first" || request.Input.Texts[1] != "second" { + t.Fatalf("embedding input mismatch: %#v", request.Input) + } +} + +func testConvertGigaChatBatchRequestItemsToJSONL(t *testing.T) { + t.Parallel() + + output, err := convertGigaChatBatchRequestItemsToJSONL(schemas.BatchEndpointChatCompletions, []schemas.BatchRequestItem{{ + CustomID: "inline-1", + Body: map[string]interface{}{ + "model": "GigaChat", + "messages": []map[string]string{ + {"role": "user", "content": "Hello from inline"}, + }, + }, + }}) + if err != nil { + t.Fatalf("convertGigaChatBatchRequestItemsToJSONL returned error: %v", err) + } + + row := decodeGigaChatBatchTestRow(t, output) + if row.ID != "inline-1" { + t.Fatalf("row id mismatch: got %q", row.ID) + } + var request GigaChatChatRequest + if err := json.Unmarshal(row.Request, &request); err != nil { + t.Fatalf("unmarshal GigaChat chat request: %v", err) + } + if len(request.Messages) != 1 || request.Messages[0].Content == nil || request.Messages[0].Content.ContentStr == nil || *request.Messages[0].Content.ContentStr != "Hello from inline" { + t.Fatalf("inline request mismatch: %#v", request) + } +} + +func testConvertGigaChatBatchInputJSONLUnsupportedEndpoint(t *testing.T) { + t.Parallel() + + input := []byte(`{"custom_id":"bad-1","method":"POST","url":"/v1/completions","body":{"model":"GigaChat","prompt":"Hello"}}` + "\n") + _, err := convertGigaChatBatchInputJSONL(schemas.BatchEndpointCompletions, input) + if err == nil { + t.Fatal("expected unsupported endpoint error, got nil") + } + if !strings.Contains(err.Error(), "line 1") || !strings.Contains(err.Error(), "do not support endpoint") { + t.Fatalf("unexpected error: %v", err) + } +} + +func testGigaChatBatchCreateTransformsFileRows(t *testing.T) { + t.Parallel() + + inputJSONL := []byte(`{"custom_id":"chat-1","method":"POST","url":"/v1/chat/completions","body":{"model":"GigaChat","messages":[{"role":"user","content":"Hello"}]}}` + "\n") + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) { + if got := request.Header.Get("Authorization"); got != "Bearer batch-token" { + t.Fatalf("authorization header mismatch: got %q", got) + } + + switch request.URL.Path { + case "/v1/files/input-file/content": + if request.Method != http.MethodGet { + t.Fatalf("file content method mismatch: got %s", request.Method) + } + w.Header().Set("Content-Type", "application/octet-stream") + _, _ = w.Write(inputJSONL) + case "/v1/batches": + if request.Method != http.MethodPost { + t.Fatalf("batch create method mismatch: got %s", request.Method) + } + if got := request.URL.Query().Get("method"); got != string(GigaChatBatchMethodChatCompletions) { + t.Fatalf("method query mismatch: got %q", got) + } + if got := request.Header.Get("Content-Type"); got != "application/octet-stream" { + t.Fatalf("content type mismatch: got %q", got) + } + body, err := io.ReadAll(request.Body) + if err != nil { + t.Fatalf("ReadAll returned error: %v", err) + } + if !bytes.HasSuffix(body, []byte("\n")) { + t.Fatalf("batch body must be JSONL with trailing newline, got %q", string(body)) + } + row := decodeGigaChatBatchTestRow(t, body) + if row.ID != "chat-1" { + t.Fatalf("row id mismatch: got %q", row.ID) + } + var chatRequest GigaChatChatRequest + if err := json.Unmarshal(row.Request, &chatRequest); err != nil { + t.Fatalf("unmarshal GigaChat batch request: %v", err) + } + if chatRequest.Model != "GigaChat" || len(chatRequest.Messages) != 1 { + t.Fatalf("unexpected chat request: %#v", chatRequest) + } + w.Header().Set("Content-Type", "application/json") + w.Header().Set("X-Request-ID", "batch-create-request-id") + _, _ = w.Write([]byte(`{"id":"batch-1","object":"batch","method":"chat_completions","status":"created","input_file_id":"input-file","completion_window":"24h","created_at":1780306293,"request_counts":{"total":1}}`)) + default: + t.Fatalf("unexpected path: %s", request.URL.Path) + } + })) + defer server.Close() + + provider := newTestGigaChatChatProvider(t, server.URL) + response, bifrostErr := provider.BatchCreate(testBifrostContext(), testGigaChatAccessTokenKey("batch-token"), &schemas.BifrostBatchCreateRequest{ + Provider: schemas.GigaChat, + InputFileID: "input-file", + Endpoint: schemas.BatchEndpointChatCompletions, + CompletionWindow: "24h", + }) + if bifrostErr != nil { + t.Fatalf("BatchCreate returned error: %v", bifrostErr) + } + if response.ID != "batch-1" || response.Status != schemas.BatchStatusValidating { + t.Fatalf("unexpected create response: %#v", response) + } + if response.Endpoint != string(schemas.BatchEndpointChatCompletions) || response.InputFileID != "input-file" || response.CompletionWindow != "24h" { + t.Fatalf("unexpected create response metadata: %#v", response) + } + if response.RequestCounts.Total != 1 { + t.Fatalf("request counts mismatch: %#v", response.RequestCounts) + } + if response.ExtraFields.Provider != schemas.GigaChat { + t.Fatalf("provider mismatch: got %q", response.ExtraFields.Provider) + } + requestID := "" + for key, value := range response.ExtraFields.ProviderResponseHeaders { + if strings.EqualFold(key, "x-request-id") { + requestID = value + break + } + } + if requestID != "batch-create-request-id" { + t.Fatalf("provider headers mismatch: %#v", response.ExtraFields.ProviderResponseHeaders) + } +} + +func testGigaChatBatchCreateUsesKeyBaseURLAndRefreshesTokenAfterUnauthorized(t *testing.T) { + t.Parallel() + + networkServer := httptest.NewServer(http.HandlerFunc(func(_ http.ResponseWriter, request *http.Request) { + t.Fatalf("network base_url server should not be used, got %s", request.URL.Path) + })) + defer networkServer.Close() + + var tokenRequests atomic.Int32 + var batchRequests atomic.Int32 + keyServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) { + switch request.URL.Path { + case "/oauth": + tokenIndex := tokenRequests.Add(1) + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"access_token":"batch-token-` + formatInt32(tokenIndex) + `","expires_at":1893456000}`)) + case "/custom-api/v1/batches": + batchIndex := batchRequests.Add(1) + wantAuthorization := "Bearer batch-token-" + formatInt32(batchIndex) + if got := request.Header.Get("Authorization"); got != wantAuthorization { + t.Fatalf("authorization header mismatch on request %d: got %q, want %q", batchIndex, got, wantAuthorization) + } + if got := request.Header.Get(gigaChatUserAgentHeader); got != gigaChatUserAgent { + t.Fatalf("user-agent mismatch: got %q", got) + } + if request.Method != http.MethodPost { + t.Fatalf("method mismatch: got %s, want POST", request.Method) + } + if got := request.URL.Query().Get("method"); got != string(GigaChatBatchMethodChatCompletions) { + t.Fatalf("method query mismatch: got %q", got) + } + if got := request.Header.Get("Content-Type"); got != "application/octet-stream" { + t.Fatalf("content type mismatch: got %q", got) + } + if batchIndex == 1 { + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusUnauthorized) + _, _ = w.Write([]byte(`{"status":401,"message":"expired token"}`)) + return + } + + body, err := io.ReadAll(request.Body) + if err != nil { + t.Fatalf("ReadAll returned error: %v", err) + } + row := decodeGigaChatBatchTestRow(t, body) + if row.ID != "inline-1" { + t.Fatalf("row id mismatch: got %q", row.ID) + } + + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"id":"batch-refreshed","object":"batch","method":"chat_completions","status":"created","completion_window":"24h","request_counts":{"total":1}}`)) + default: + t.Fatalf("unexpected path: %s", request.URL.Path) + } + })) + defer keyServer.Close() + + provider := newTestGigaChatChatProvider(t, networkServer.URL) + key := schemas.Key{ + GigaChatKeyConfig: &schemas.GigaChatKeyConfig{ + Credentials: schemas.NewSecretVar("test-credentials"), + AuthURL: keyServer.URL + "/oauth", + BaseURL: keyServer.URL + "/custom-api", + }, + } + + response, bifrostErr := provider.BatchCreate(testBifrostContext(), key, &schemas.BifrostBatchCreateRequest{ + Provider: schemas.GigaChat, + Endpoint: schemas.BatchEndpointChatCompletions, + CompletionWindow: "24h", + Requests: []schemas.BatchRequestItem{{ + CustomID: "inline-1", + Body: map[string]interface{}{ + "model": "GigaChat", + "messages": []map[string]string{ + {"role": "user", "content": "Hello"}, + }, + }, + }}, + }) + if bifrostErr != nil { + t.Fatalf("BatchCreate returned error: %v", bifrostErr) + } + if response.ID != "batch-refreshed" || response.Status != schemas.BatchStatusValidating { + t.Fatalf("unexpected response: %#v", response) + } + if tokenRequests.Load() != 2 { + t.Fatalf("token request count mismatch: got %d, want 2", tokenRequests.Load()) + } + if batchRequests.Load() != 2 { + t.Fatalf("batch request count mismatch: got %d, want 2", batchRequests.Load()) + } +} + +func testGigaChatBatchListParsesWrapper(t *testing.T) { + t.Parallel() + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) { + if request.URL.Path != "/v1/batches" { + t.Fatalf("path mismatch: got %s", request.URL.Path) + } + if request.Method != http.MethodGet { + t.Fatalf("method mismatch: got %s", request.Method) + } + if request.URL.RawQuery != "" { + t.Fatalf("unexpected query: %s", request.URL.RawQuery) + } + if got := request.Header.Get("Authorization"); got != "Bearer batch-list-token" { + t.Fatalf("authorization header mismatch: got %q", got) + } + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"data":[{"id":"batch-1","object":"batch","method":"chat_completions","status":"in_progress","created_at":1780306293,"request_counts":{"total":2,"completed":1}},{"id":"batch-2","object":"batch","method":"embedder","status":"completed","created_at":1780306294,"request_counts":{"total":1,"completed":1}}]}`)) + })) + defer server.Close() + + provider := newTestGigaChatChatProvider(t, server.URL) + response, bifrostErr := provider.BatchList(testBifrostContext(), []schemas.Key{testGigaChatAccessTokenKey("batch-list-token")}, &schemas.BifrostBatchListRequest{Provider: schemas.GigaChat}) + if bifrostErr != nil { + t.Fatalf("BatchList returned error: %v", bifrostErr) + } + if response.Object != "list" || len(response.Data) != 2 { + t.Fatalf("unexpected list response: %#v", response) + } + if response.Data[0].Status != schemas.BatchStatusInProgress || response.Data[0].Endpoint != string(schemas.BatchEndpointChatCompletions) { + t.Fatalf("first batch mismatch: %#v", response.Data[0]) + } + if response.Data[1].Status != schemas.BatchStatusCompleted || response.Data[1].Endpoint != string(schemas.BatchEndpointEmbeddings) { + t.Fatalf("second batch mismatch: %#v", response.Data[1]) + } + if response.FirstID == nil || *response.FirstID != "batch-1" || response.LastID == nil || *response.LastID != "batch-2" { + t.Fatalf("list ids mismatch: first=%v last=%v", response.FirstID, response.LastID) + } +} + +func testGigaChatBatchListParsesEmptyRootArray(t *testing.T) { + t.Parallel() + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) { + if request.URL.Path != "/v1/batches" { + t.Fatalf("path mismatch: got %s", request.URL.Path) + } + if request.Method != http.MethodGet { + t.Fatalf("method mismatch: got %s", request.Method) + } + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`[]`)) + })) + defer server.Close() + + provider := newTestGigaChatChatProvider(t, server.URL) + response, bifrostErr := provider.BatchList(testBifrostContext(), []schemas.Key{testGigaChatAccessTokenKey("batch-list-token")}, &schemas.BifrostBatchListRequest{Provider: schemas.GigaChat}) + if bifrostErr != nil { + t.Fatalf("BatchList returned error: %v", bifrostErr) + } + if response.Object != "list" || len(response.Data) != 0 { + t.Fatalf("unexpected list response: %#v", response) + } + if response.FirstID != nil || response.LastID != nil { + t.Fatalf("empty list should not set cursors: first=%v last=%v", response.FirstID, response.LastID) + } +} + +func testGigaChatBatchListPaginatesLocally(t *testing.T) { + t.Parallel() + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) { + if request.URL.Path != "/v1/batches" { + t.Fatalf("path mismatch: got %s", request.URL.Path) + } + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"data":[{"id":"batch-1","object":"batch","status":"in_progress"},{"id":"batch-2","object":"batch","status":"completed"}]}`)) + })) + defer server.Close() + + provider := newTestGigaChatChatProvider(t, server.URL) + request := &schemas.BifrostBatchListRequest{Provider: schemas.GigaChat, Limit: 1} + first, bifrostErr := provider.BatchList(testBifrostContext(), []schemas.Key{testGigaChatAccessTokenKey("batch-list-token")}, request) + if bifrostErr != nil { + t.Fatalf("first BatchList returned error: %v", bifrostErr) + } + if len(first.Data) != 1 || first.Data[0].ID != "batch-1" || !first.HasMore || first.NextCursor == nil { + t.Fatalf("unexpected first page: %#v", first) + } + + request.After = first.NextCursor + second, bifrostErr := provider.BatchList(testBifrostContext(), []schemas.Key{testGigaChatAccessTokenKey("batch-list-token")}, request) + if bifrostErr != nil { + t.Fatalf("second BatchList returned error: %v", bifrostErr) + } + if len(second.Data) != 1 || second.Data[0].ID != "batch-2" || second.HasMore || second.NextCursor != nil { + t.Fatalf("unexpected second page: %#v", second) + } +} + +func testGigaChatBatchRetrieveParsesSingleObject(t *testing.T) { + t.Parallel() + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) { + if request.URL.Path != "/v1/batches" { + t.Fatalf("path mismatch: got %s", request.URL.Path) + } + if request.Method != http.MethodGet { + t.Fatalf("method mismatch: got %s", request.Method) + } + if got := request.URL.Query().Get("batch_id"); got != "batch-1" { + t.Fatalf("batch_id query mismatch: got %q", got) + } + if got := request.Header.Get("Authorization"); got != "Bearer batch-retrieve-token" { + t.Fatalf("authorization header mismatch: got %q", got) + } + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"id":"batch-1","object":"batch","method":"embedder","status":"completed","created_at":1780306293,"completed_at":1780306393,"output_file_id":"output-file","error_file_id":"error-file","request_counts":{"total":3,"completed":2,"failed":1}}`)) + })) + defer server.Close() + + provider := newTestGigaChatChatProvider(t, server.URL) + response, bifrostErr := provider.BatchRetrieve(testBifrostContext(), []schemas.Key{testGigaChatAccessTokenKey("batch-retrieve-token")}, &schemas.BifrostBatchRetrieveRequest{ + Provider: schemas.GigaChat, + BatchID: "batch-1", + }) + if bifrostErr != nil { + t.Fatalf("BatchRetrieve returned error: %v", bifrostErr) + } + if response.ID != "batch-1" || response.Status != schemas.BatchStatusCompleted || response.Endpoint != string(schemas.BatchEndpointEmbeddings) { + t.Fatalf("unexpected retrieve response: %#v", response) + } + if response.OutputFileID == nil || *response.OutputFileID != "output-file" || response.ErrorFileID == nil || *response.ErrorFileID != "error-file" { + t.Fatalf("file ids mismatch: output=%v error=%v", response.OutputFileID, response.ErrorFileID) + } + if response.RequestCounts.Total != 3 || response.RequestCounts.Completed != 2 || response.RequestCounts.Failed != 1 { + t.Fatalf("request counts mismatch: %#v", response.RequestCounts) + } +} + +func testGigaChatBatchRetrieveMapsResultFileID(t *testing.T) { + t.Parallel() + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) { + if request.URL.Path != "/v1/batches" { + t.Fatalf("path mismatch: got %s", request.URL.Path) + } + if request.Method != http.MethodGet { + t.Fatalf("method mismatch: got %s", request.Method) + } + if got := request.URL.Query().Get("batch_id"); got != "batch-1" { + t.Fatalf("batch_id query mismatch: got %q", got) + } + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"id":"batch-1","object":"batch","method":"chat_completions","status":"completed","result_file_id":"result-file","request_counts":{"total":1,"completed":1}}`)) + })) + defer server.Close() + + provider := newTestGigaChatChatProvider(t, server.URL) + response, bifrostErr := provider.BatchRetrieve(testBifrostContext(), []schemas.Key{testGigaChatAccessTokenKey("batch-token")}, &schemas.BifrostBatchRetrieveRequest{ + Provider: schemas.GigaChat, + BatchID: "batch-1", + }) + if bifrostErr != nil { + t.Fatalf("BatchRetrieve returned error: %v", bifrostErr) + } + if response.OutputFileID == nil || *response.OutputFileID != "result-file" { + t.Fatalf("result_file_id was not mapped to output_file_id: %#v", response.OutputFileID) + } +} + +func testGigaChatBatchPreservesResponsesEndpoint(t *testing.T) { + t.Parallel() + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) { + if request.URL.Path != "/v1/batches" { + t.Fatalf("path mismatch: got %s", request.URL.Path) + } + if got := request.Header.Get("Authorization"); got != "Bearer batch-responses-token" { + t.Fatalf("authorization header mismatch: got %q", got) + } + w.Header().Set("Content-Type", "application/json") + + switch request.Method { + case http.MethodPost: + if got := request.URL.Query().Get("method"); got != string(GigaChatBatchMethodResponses) { + t.Fatalf("method query mismatch: got %q, want %q", got, GigaChatBatchMethodResponses) + } + body, err := io.ReadAll(request.Body) + if err != nil { + t.Fatalf("ReadAll returned error: %v", err) + } + row := decodeGigaChatBatchTestRow(t, body) + if row.ID != "responses-1" { + t.Fatalf("row id mismatch: got %q", row.ID) + } + _, _ = w.Write([]byte(`{"id":"batch-responses","object":"batch","method":"responses","status":"created","completion_window":"24h","request_counts":{"total":1}}`)) + case http.MethodGet: + if got := request.URL.Query().Get("batch_id"); got == "" { + _, _ = w.Write([]byte(`{"data":[{"id":"batch-responses","object":"batch","method":"responses","status":"in_progress","created_at":1780306293,"request_counts":{"total":1}}]}`)) + } else if got == "batch-responses" { + _, _ = w.Write([]byte(`{"id":"batch-responses","object":"batch","method":"responses","status":"completed","created_at":1780306293,"request_counts":{"total":1,"completed":1}}`)) + } else { + t.Fatalf("batch_id query mismatch: got %q", got) + } + default: + t.Fatalf("method mismatch: got %s", request.Method) + } + })) + defer server.Close() + + provider := newTestGigaChatChatProvider(t, server.URL) + key := testGigaChatAccessTokenKey("batch-responses-token") + createResponse, bifrostErr := provider.BatchCreate(testBifrostContext(), key, &schemas.BifrostBatchCreateRequest{ + Provider: schemas.GigaChat, + Endpoint: schemas.BatchEndpointResponses, + CompletionWindow: "24h", + Requests: []schemas.BatchRequestItem{{ + CustomID: "responses-1", + Body: map[string]interface{}{ + "model": "GigaChat-2", + "input": "Summarize this", + }, + }}, + }) + if bifrostErr != nil { + t.Fatalf("BatchCreate returned error: %v", bifrostErr) + } + if createResponse.Endpoint != string(schemas.BatchEndpointResponses) { + t.Fatalf("create endpoint mismatch: got %q", createResponse.Endpoint) + } + + listResponse, bifrostErr := provider.BatchList(testBifrostContext(), []schemas.Key{key}, &schemas.BifrostBatchListRequest{Provider: schemas.GigaChat}) + if bifrostErr != nil { + t.Fatalf("BatchList returned error: %v", bifrostErr) + } + if len(listResponse.Data) != 1 || listResponse.Data[0].Endpoint != string(schemas.BatchEndpointResponses) { + t.Fatalf("list did not preserve responses endpoint: %#v", listResponse.Data) + } + + retrieveResponse, bifrostErr := provider.BatchRetrieve(testBifrostContext(), []schemas.Key{key}, &schemas.BifrostBatchRetrieveRequest{ + Provider: schemas.GigaChat, + BatchID: "batch-responses", + }) + if bifrostErr != nil { + t.Fatalf("BatchRetrieve returned error: %v", bifrostErr) + } + if retrieveResponse.Endpoint != string(schemas.BatchEndpointResponses) { + t.Fatalf("retrieve endpoint mismatch: got %q", retrieveResponse.Endpoint) + } +} + +func testGigaChatBatchResultsDownloadsOutputFile(t *testing.T) { + t.Parallel() + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) { + if got := request.Header.Get("Authorization"); got != "Bearer batch-results-token" { + t.Fatalf("authorization header mismatch: got %q", got) + } + switch request.URL.Path { + case "/v1/batches": + if request.Method != http.MethodGet { + t.Fatalf("method mismatch: got %s", request.Method) + } + if got := request.URL.Query().Get("batch_id"); got != "batch-1" { + t.Fatalf("batch_id query mismatch: got %q", got) + } + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"id":"batch-1","object":"batch","method":"chat_completions","status":"completed","result_file_id":"result-file","request_counts":{"total":1,"completed":1}}`)) + case "/v1/files/result-file/content": + if request.Method != http.MethodGet { + t.Fatalf("method mismatch: got %s", request.Method) + } + w.Header().Set("Content-Type", "application/octet-stream") + _, _ = w.Write([]byte(`{"id":"row-1","response":{"status_code":200,"request_id":"req-1","body":{"id":"chatcmpl-1","object":"chat.completion"}}}` + "\n")) + default: + t.Fatalf("unexpected path: %s", request.URL.Path) + } + })) + defer server.Close() + + provider := newTestGigaChatChatProvider(t, server.URL) + response, bifrostErr := provider.BatchResults(testBifrostContext(), []schemas.Key{testGigaChatAccessTokenKey("batch-results-token")}, &schemas.BifrostBatchResultsRequest{ + Provider: schemas.GigaChat, + BatchID: "batch-1", + }) + if bifrostErr != nil { + t.Fatalf("BatchResults returned error: %v", bifrostErr) + } + if response.BatchID != "batch-1" || len(response.Results) != 1 { + t.Fatalf("unexpected batch results response: %#v", response) + } + result := response.Results[0] + if result.CustomID != "row-1" || result.Response == nil || result.Response.StatusCode != 200 || result.Response.RequestID != "req-1" { + t.Fatalf("unexpected result item: %#v", result) + } +} + +func testGigaChatBatchResultsWithoutOutputFileID(t *testing.T) { + t.Parallel() + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) { + if request.URL.Path != "/v1/batches" { + t.Fatalf("path mismatch: got %s", request.URL.Path) + } + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"id":"batch-1","object":"batch","method":"chat_completions","status":"in_progress","request_counts":{"total":1}}`)) + })) + defer server.Close() + + provider := newTestGigaChatChatProvider(t, server.URL) + response, bifrostErr := provider.BatchResults(testBifrostContext(), []schemas.Key{testGigaChatAccessTokenKey("batch-results-token")}, &schemas.BifrostBatchResultsRequest{ + Provider: schemas.GigaChat, + BatchID: "batch-1", + }) + if response != nil { + t.Fatalf("expected nil response, got %#v", response) + } + if bifrostErr == nil || !strings.Contains(bifrostErr.Error.Message, "did not return output_file_id or result_file_id") { + t.Fatalf("unexpected error: %v", bifrostErr) + } +} + +func testGigaChatBatchOutputFileRejectsEmptyKeys(t *testing.T) { + t.Parallel() + + provider, err := NewGigaChatProvider(&schemas.ProviderConfig{}, nil) + if err != nil { + t.Fatalf("NewGigaChatProvider returned error: %v", err) + } + + response, bifrostErr := provider.readGigaChatBatchOutputFile(testBifrostContext(), nil, &schemas.BifrostBatchResultsRequest{ + Provider: schemas.GigaChat, + BatchID: "batch-1", + }, "result-file") + if response != nil { + t.Fatalf("expected nil response, got %#v", response) + } + if bifrostErr == nil || bifrostErr.Error == nil || !strings.Contains(bifrostErr.Error.Message, "no keys available") { + t.Fatalf("unexpected error: %v", bifrostErr) + } +} + +func testGigaChatBatchCreateUnsupportedEndpoint(t *testing.T) { + t.Parallel() + + provider, err := NewGigaChatProvider(&schemas.ProviderConfig{}, nil) + if err != nil { + t.Fatalf("NewGigaChatProvider returned error: %v", err) + } + + response, bifrostErr := provider.BatchCreate(testBifrostContext(), testGigaChatAccessTokenKey("batch-token"), &schemas.BifrostBatchCreateRequest{ + Provider: schemas.GigaChat, + Endpoint: schemas.BatchEndpointCompletions, + Requests: []schemas.BatchRequestItem{{ + CustomID: "bad-1", + Body: map[string]interface{}{ + "model": "GigaChat", + "prompt": "Hello", + }, + }}, + }) + if response != nil { + t.Fatalf("expected nil response, got %#v", response) + } + if bifrostErr == nil || !strings.Contains(bifrostErr.Error.Message, "do not support endpoint") { + t.Fatalf("unexpected error: %v", bifrostErr) + } +} + +func testGigaChatBatchCreateUnsupportedCompletionWindow(t *testing.T) { + t.Parallel() + + provider, err := NewGigaChatProvider(&schemas.ProviderConfig{}, nil) + if err != nil { + t.Fatalf("NewGigaChatProvider returned error: %v", err) + } + + response, bifrostErr := provider.BatchCreate(testBifrostContext(), testGigaChatAccessTokenKey("batch-token"), &schemas.BifrostBatchCreateRequest{ + Provider: schemas.GigaChat, + Endpoint: schemas.BatchEndpointChatCompletions, + CompletionWindow: "1h", + Requests: []schemas.BatchRequestItem{{ + CustomID: "chat-1", + Body: map[string]interface{}{ + "model": "GigaChat", + "messages": []map[string]string{ + {"role": "user", "content": "Hello"}, + }, + }, + }}, + }) + if response != nil { + t.Fatalf("expected nil response, got %#v", response) + } + if bifrostErr == nil || !strings.Contains(bifrostErr.Error.Message, "completion_window=24h only") { + t.Fatalf("unexpected error: %v", bifrostErr) + } +} + +func decodeGigaChatBatchTestRow(t *testing.T, output []byte) GigaChatBatchInputRow { + t.Helper() + + lines := strings.Split(strings.TrimSpace(string(output)), "\n") + if len(lines) != 1 { + t.Fatalf("got %d JSONL lines, want 1: %q", len(lines), string(output)) + } + var row GigaChatBatchInputRow + if err := json.Unmarshal([]byte(lines[0]), &row); err != nil { + t.Fatalf("unmarshal GigaChat batch row: %v", err) + } + if len(row.Request) == 0 { + t.Fatalf("row request is empty: %#v", row) + } + return row +} diff --git a/core/providers/gigachat/chat_test.go b/core/providers/gigachat/chat_test.go new file mode 100644 index 00000000000..2fa519d34a0 --- /dev/null +++ b/core/providers/gigachat/chat_test.go @@ -0,0 +1,1870 @@ +package gigachat + +import ( + "context" + "encoding/json" + "fmt" + "io" + "net" + "net/http" + "net/http/httptest" + "strings" + "sync/atomic" + "testing" + "time" + + providerUtils "github.com/maximhq/bifrost/core/providers/utils" + "github.com/maximhq/bifrost/core/schemas" +) + +func TestGigaChatChatCompletion(t *testing.T) { + testGigaChatChatCompletion(t) +} + +func TestGigaChatChatCompletionFileDataDecoding(t *testing.T) { + t.Parallel() + + t.Run("RawTextFileData", func(t *testing.T) { + t.Parallel() + + filename := "note.txt" + fileType := "text/plain" + fileData := "test" + upload, err := gigaChatChatFileUpload(3, &schemas.ChatInputFile{ + Filename: &filename, + FileData: &fileData, + FileType: &fileType, + }) + if err != nil { + t.Fatalf("gigaChatChatFileUpload returned error: %v", err) + } + if string(upload.file) != "test" { + t.Fatalf("raw text file_data was not preserved: %q", string(upload.file)) + } + if upload.filename != "note.txt" || upload.contentType != "text/plain" { + t.Fatalf("upload metadata mismatch: %#v", upload) + } + }) + + t.Run("TextDataURLBase64", func(t *testing.T) { + t.Parallel() + + filename := "note.txt" + fileData := "data:text/plain;base64,dGVzdA==" + upload, err := gigaChatChatFileUpload(4, &schemas.ChatInputFile{ + Filename: &filename, + FileData: &fileData, + }) + if err != nil { + t.Fatalf("gigaChatChatFileUpload returned error: %v", err) + } + if string(upload.file) != "test" { + t.Fatalf("base64 data URL was not decoded: %q", string(upload.file)) + } + if !strings.HasPrefix(upload.contentType, "text/plain") { + t.Fatalf("content type mismatch: %q", upload.contentType) + } + }) + + t.Run("NonTextInvalidBase64IncludesBlockIndex", func(t *testing.T) { + t.Parallel() + + filename := "document.pdf" + fileType := "application/pdf" + fileData := "not-base64!" + _, err := gigaChatChatFileUpload(7, &schemas.ChatInputFile{ + Filename: &filename, + FileData: &fileData, + FileType: &fileType, + }) + if err == nil { + t.Fatal("expected non-text invalid base64 to fail") + } + if !strings.Contains(err.Error(), "content block 7") || !strings.Contains(err.Error(), "file_data must be a base64 data URL or base64-encoded content") { + t.Fatalf("unexpected error: %v", err) + } + }) +} + +func TestGigaChatChatCompletionStreamToolCallIndex(t *testing.T) { + t.Parallel() + + functionsStateID := "call-weather" + response := ToBifrostChatStreamResponse(schemas.GigaChat, &GigaChatChatStreamResponse{ + Model: "GigaChat", + Choices: []GigaChatChatStreamChoice{{ + Index: 2, + Delta: &GigaChatChatStreamDelta{ + FunctionCall: &GigaChatFunctionCall{ + Name: "get_weather", + Arguments: json.RawMessage(`{"city":"Moscow"}`), + }, + FunctionsStateID: &functionsStateID, + }, + }}, + }) + if response == nil || len(response.Choices) != 1 || response.Choices[0].ChatStreamResponseChoice == nil { + t.Fatalf("stream response mismatch: %#v", response) + } + choice := response.Choices[0] + if choice.Index != 2 { + t.Fatalf("choice index mismatch: got %d, want 2", choice.Index) + } + delta := choice.ChatStreamResponseChoice.Delta + if delta == nil || len(delta.ToolCalls) != 1 { + t.Fatalf("tool call delta mismatch: %#v", delta) + } + toolCall := delta.ToolCalls[0] + if toolCall.Index != 0 { + t.Fatalf("tool call index mismatch: got %d, want 0", toolCall.Index) + } + if toolCall.ID == nil || *toolCall.ID != functionsStateID { + t.Fatalf("tool call id mismatch: %#v", toolCall.ID) + } + if toolCall.Function.Name == nil || *toolCall.Function.Name != "get_weather" || toolCall.Function.Arguments != `{"city":"Moscow"}` { + t.Fatalf("tool call function mismatch: %#v", toolCall.Function) + } +} + +func testGigaChatChatCompletion(t *testing.T) { + t.Parallel() + + t.Run("ConverterMapsRequest", testGigaChatChatConverterMapsRequest) + t.Run("ConverterMapsOpenAIJSONSchemaResponseFormat", testGigaChatChatConverterMapsOpenAIJSONSchemaResponseFormat) + t.Run("ConverterMapsGigaChatJSONSchemaResponseFormat", testGigaChatChatConverterMapsGigaChatJSONSchemaResponseFormat) + t.Run("ConverterPreservesAssistantReasoningContent", testGigaChatChatConverterPreservesAssistantReasoningContent) + t.Run("ConverterMapsFileAttachments", testGigaChatChatConverterMapsFileAttachments) + t.Run("ExecutesWithOAuthTokenAndExtraParams", testGigaChatChatCompletionExecutesWithOAuthTokenAndExtraParams) + t.Run("ExecutesWithMTLSClientCertificate", testGigaChatChatCompletionExecutesWithMTLSClientCertificate) + t.Run("UploadsInlineImageAttachment", testGigaChatChatCompletionUploadsInlineImageAttachment) + t.Run("UploadsInlineFileAttachment", testGigaChatChatCompletionUploadsInlineFileAttachment) + t.Run("ReusesUploadedAttachmentAfterBackendError", testGigaChatChatCompletionReusesUploadedAttachmentAfterBackendError) + t.Run("ReusesSuccessfulAttachmentAfterPartialUploadFailure", testGigaChatChatCompletionReusesSuccessfulAttachmentAfterPartialUploadFailure) + t.Run("DoesNotReuseUploadedAttachmentAcrossIndependentRequests", testGigaChatChatCompletionDoesNotReuseUploadedAttachmentAcrossIndependentRequests) + t.Run("DoesNotCacheFailedAttachmentUpload", testGigaChatChatCompletionDoesNotCacheFailedAttachmentUpload) + t.Run("RejectsUnsupportedTools", testGigaChatChatCompletionRejectsUnsupportedTools) + t.Run("RejectsUnsupportedResponseFormat", testGigaChatChatCompletionRejectsUnsupportedResponseFormat) + t.Run("MapsProviderErrors", testGigaChatChatCompletionMapsProviderErrors) + t.Run("RefreshesTokenAfterUnauthorized", testGigaChatChatCompletionRefreshesTokenAfterUnauthorized) + t.Run("DoesNotDoubleExchangeExpiredTokenOnRefresh", testGigaChatChatCompletionDoesNotDoubleExchangeExpiredTokenOnRefresh) + t.Run("StreamsSSEChunks", testGigaChatChatCompletionStreamsSSEChunks) + t.Run("FinalizesLargeResponsePassthrough", testGigaChatChatCompletionStreamFinalizesLargeResponsePassthrough) + t.Run("MapsStreamingProviderErrors", testGigaChatChatCompletionMapsStreamingProviderErrors) + t.Run("MapsStreamingErrorEvents", testGigaChatChatCompletionMapsStreamingErrorEvents) + t.Run("RefreshesStreamingTokenAfterUnauthorized", testGigaChatChatCompletionStreamRefreshesTokenAfterUnauthorized) + t.Run("HandlesStreamingContextCancellation", testGigaChatChatCompletionStreamHandlesContextCancellation) +} + +func testGigaChatChatConverterMapsRequest(t *testing.T) { + t.Parallel() + + maxTokens := 512 + temperature := 0.2 + topP := 0.8 + n := 1 + reasoningEffort := "high" + text := "hello" + request := &schemas.BifrostChatRequest{ + Model: "GigaChat", + Input: []schemas.ChatMessage{ + { + Role: schemas.ChatMessageRoleUser, + Content: &schemas.ChatMessageContent{ContentStr: &text}, + }, + }, + Params: &schemas.ChatParameters{ + MaxCompletionTokens: &maxTokens, + Temperature: &temperature, + TopP: &topP, + N: &n, + Stop: []string{"stop"}, + Reasoning: &schemas.ChatReasoning{Effort: &reasoningEffort}, + ExtraParams: map[string]interface{}{ + "profanity_check": false, + }, + }, + } + + gigaChatReq, err := ToGigaChatChatRequest(testBifrostContext(), request) + if err != nil { + t.Fatalf("ToGigaChatChatRequest returned error: %v", err) + } + if gigaChatReq.MaxTokens == nil || *gigaChatReq.MaxTokens != maxTokens { + t.Fatalf("max_tokens mismatch: got %#v, want %d", gigaChatReq.MaxTokens, maxTokens) + } + if gigaChatReq.Stream == nil || *gigaChatReq.Stream { + t.Fatalf("stream mismatch: got %#v, want false", gigaChatReq.Stream) + } + if got := gigaChatReq.GetExtraParams()["profanity_check"]; got != false { + t.Fatalf("extra param mismatch: got %#v", got) + } + if gigaChatReq.ReasoningEffort == nil || *gigaChatReq.ReasoningEffort != reasoningEffort { + t.Fatalf("reasoning_effort mismatch: got %#v, want %q", gigaChatReq.ReasoningEffort, reasoningEffort) + } + + body, err := json.Marshal(gigaChatReq) + if err != nil { + t.Fatalf("failed to marshal request: %v", err) + } + if strings.Contains(string(body), "max_completion_tokens") { + t.Fatalf("request body should use max_tokens, got %s", body) + } + if !strings.Contains(string(body), `"max_tokens":512`) { + t.Fatalf("request body missing max_tokens: %s", body) + } + if strings.Contains(string(body), `"reasoning":`) { + t.Fatalf("request body should not send reasoning object: %s", body) + } + if !strings.Contains(string(body), `"reasoning_effort":"high"`) { + t.Fatalf("request body missing reasoning_effort: %s", body) + } +} + +func testGigaChatChatConverterMapsOpenAIJSONSchemaResponseFormat(t *testing.T) { + t.Parallel() + + strict := true + formatName := "MathAnswer" + formatDescription := "Math answer schema." + responseFormat := interface{}(map[string]interface{}{ + "type": "json_schema", + "json_schema": map[string]interface{}{ + "name": formatName, + "description": formatDescription, + "strict": strict, + "schema": map[string]interface{}{ + "type": "object", + "properties": map[string]interface{}{ + "steps": map[string]interface{}{"type": "array", "items": map[string]interface{}{"type": "string"}}, + "final_answer": map[string]interface{}{"type": "string"}, + }, + "required": []interface{}{"steps", "final_answer"}, + }, + }, + }) + + request := testGigaChatChatRequest() + request.Params.ResponseFormat = &responseFormat + + gigaChatReq, err := ToGigaChatChatRequest(testBifrostContext(), request) + if err != nil { + t.Fatalf("ToGigaChatChatRequest returned error: %v", err) + } + assertGigaChatJSONSchemaResponseFormat(t, gigaChatReq.ResponseFormat, formatName, formatDescription, strict) + + body, err := json.Marshal(gigaChatReq) + if err != nil { + t.Fatalf("failed to marshal request: %v", err) + } + if strings.Contains(string(body), `"json_schema":`) { + t.Fatalf("GigaChat response_format should use schema, not OpenAI json_schema wrapper: %s", body) + } + if !strings.Contains(string(body), `"response_format"`) || !strings.Contains(string(body), `"schema"`) { + t.Fatalf("request body missing response_format schema: %s", body) + } +} + +func testGigaChatChatConverterMapsGigaChatJSONSchemaResponseFormat(t *testing.T) { + t.Parallel() + + responseFormat := interface{}(schemas.NewOrderedMapFromPairs( + schemas.KV("type", "json_schema"), + schemas.KV("schema", schemas.NewOrderedMapFromPairs( + schemas.KV("type", "object"), + schemas.KV("properties", schemas.NewOrderedMapFromPairs( + schemas.KV("status", schemas.NewOrderedMapFromPairs(schemas.KV("type", "string"))), + )), + schemas.KV("required", []interface{}{"status"}), + )), + schemas.KV("strict", true), + )) + + request := testGigaChatChatRequest() + request.Params.ResponseFormat = &responseFormat + + gigaChatReq, err := ToGigaChatChatRequest(testBifrostContext(), request) + if err != nil { + t.Fatalf("ToGigaChatChatRequest returned error: %v", err) + } + assertGigaChatJSONSchemaResponseFormat(t, gigaChatReq.ResponseFormat, "", "", true) +} + +func testGigaChatChatConverterPreservesAssistantReasoningContent(t *testing.T) { + t.Parallel() + + userText := "question" + answerText := "answer" + reasoning := "model reasoning" + request := &schemas.BifrostChatRequest{ + Model: "GigaChat", + Input: []schemas.ChatMessage{ + { + Role: schemas.ChatMessageRoleUser, + Content: &schemas.ChatMessageContent{ContentStr: &userText}, + }, + { + Role: schemas.ChatMessageRoleAssistant, + Content: &schemas.ChatMessageContent{ContentStr: &answerText}, + ChatAssistantMessage: &schemas.ChatAssistantMessage{ + Reasoning: &reasoning, + }, + }, + }, + } + + gigaChatReq, err := ToGigaChatChatRequest(testBifrostContext(), request) + if err != nil { + t.Fatalf("ToGigaChatChatRequest returned error: %v", err) + } + if len(gigaChatReq.Messages) != 2 { + t.Fatalf("message count mismatch: got %d", len(gigaChatReq.Messages)) + } + assistant := gigaChatReq.Messages[1] + if assistant.Reasoning == nil || *assistant.Reasoning != reasoning { + t.Fatalf("assistant reasoning_content mismatch: %#v", assistant.Reasoning) + } + + body, err := json.Marshal(gigaChatReq) + if err != nil { + t.Fatalf("failed to marshal request: %v", err) + } + if !strings.Contains(string(body), `"reasoning_content":"model reasoning"`) { + t.Fatalf("request body missing reasoning_content: %s", body) + } + if strings.Contains(string(body), `"reasoning":"model reasoning"`) { + t.Fatalf("request body should use reasoning_content, got %s", body) + } +} + +func testGigaChatChatConverterMapsFileAttachments(t *testing.T) { + t.Parallel() + + prompt := "Кратко перескажи документ" + fileID := "file-document" + filename := "document.pdf" + fileType := "application/pdf" + request := &schemas.BifrostChatRequest{ + Model: "GigaChat", + Input: []schemas.ChatMessage{ + { + Role: schemas.ChatMessageRoleUser, + Content: &schemas.ChatMessageContent{ + ContentBlocks: []schemas.ChatContentBlock{ + {Type: schemas.ChatContentBlockTypeText, Text: &prompt}, + { + Type: schemas.ChatContentBlockTypeFile, + File: &schemas.ChatInputFile{ + FileID: &fileID, + Filename: &filename, + FileType: &fileType, + }, + }, + }, + }, + }, + }, + } + + gigaChatReq, err := ToGigaChatChatRequest(testBifrostContext(), request) + if err != nil { + t.Fatalf("ToGigaChatChatRequest returned error: %v", err) + } + if len(gigaChatReq.Messages) != 1 { + t.Fatalf("message count mismatch: got %d", len(gigaChatReq.Messages)) + } + message := gigaChatReq.Messages[0] + if message.Content == nil || message.Content.ContentStr == nil || *message.Content.ContentStr != prompt { + t.Fatalf("content mismatch: %#v", message.Content) + } + if len(message.Attachments) != 1 || message.Attachments[0] != fileID { + t.Fatalf("attachments mismatch: %#v", message.Attachments) + } + if gigaChatReq.FunctionCall != "auto" { + t.Fatalf("function_call mismatch: got %#v, want auto", gigaChatReq.FunctionCall) + } + + body, err := json.Marshal(gigaChatReq) + if err != nil { + t.Fatalf("failed to marshal request: %v", err) + } + bodyStr := string(body) + if !strings.Contains(bodyStr, `"attachments":["file-document"]`) { + t.Fatalf("request body missing attachments: %s", body) + } + if strings.Contains(bodyStr, "file_data") || strings.Contains(bodyStr, "file_id") { + t.Fatalf("request body should not include OpenAI file content block fields: %s", body) + } +} + +func testGigaChatChatCompletionExecutesWithOAuthTokenAndExtraParams(t *testing.T) { + t.Parallel() + + var tokenRequests atomic.Int32 + var chatRequests atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) { + switch request.URL.Path { + case "/oauth": + tokenRequests.Add(1) + if got := request.Header.Get("Authorization"); got != "Basic super-secret-credentials" { + t.Fatalf("token authorization header mismatch: got %q", got) + } + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"access_token":"chat-access-token","expires_at":1893456000}`)) + case "/v1/chat/completions": + chatRequests.Add(1) + if got := request.Header.Get("Authorization"); got != "Bearer chat-access-token" { + t.Fatalf("chat authorization header mismatch: got %q", got) + } + if strings.Contains(request.Header.Get("Authorization"), "super-secret-credentials") { + t.Fatal("chat request leaked OAuth credentials") + } + assertGigaChatChatRequestBody(t, request) + w.Header().Set("Content-Type", "application/json") + w.Header().Set("X-Request-ID", "chat-request-id") + _, _ = w.Write([]byte(`{ + "id":"chatcmpl-test", + "choices":[{"index":0,"message":{"role":"assistant","content":"Здравствуйте","reasoning_content":"Думаю"},"finish_reason":"stop"}], + "created":1700000000, + "model":"GigaChat", + "object":"chat.completion", + "usage":{"prompt_tokens":7,"completion_tokens":3,"total_tokens":10,"precached_prompt_tokens":2} + }`)) + default: + t.Fatalf("unexpected path: %s", request.URL.Path) + } + })) + defer server.Close() + + provider := newTestGigaChatChatProvider(t, server.URL) + provider.sendBackRawRequest = true + provider.sendBackRawResponse = true + ctx := testBifrostContext() + ctx.SetValue(schemas.BifrostContextKeyCaptureRawRequest, true) + ctx.SetValue(schemas.BifrostContextKeyCaptureRawResponse, true) + ctx.SetValue(schemas.BifrostContextKeyPassthroughExtraParams, true) + + response, bifrostErr := provider.ChatCompletion(ctx, testGigaChatOAuthKey(server.URL+"/oauth", "", "super-secret-credentials"), testGigaChatChatRequest()) + if bifrostErr != nil { + t.Fatalf("ChatCompletion returned error: %v", bifrostErr) + } + if tokenRequests.Load() != 1 { + t.Fatalf("token request count mismatch: got %d, want 1", tokenRequests.Load()) + } + if chatRequests.Load() != 1 { + t.Fatalf("chat request count mismatch: got %d, want 1", chatRequests.Load()) + } + if response.ID != "chatcmpl-test" { + t.Fatalf("response id mismatch: got %q", response.ID) + } + if response.ExtraFields.Provider != schemas.GigaChat { + t.Fatalf("provider mismatch: got %q, want %q", response.ExtraFields.Provider, schemas.GigaChat) + } + if response.Usage == nil || response.Usage.TotalTokens != 10 { + t.Fatalf("usage mismatch: %#v", response.Usage) + } + if response.Usage.PromptTokensDetails == nil || response.Usage.PromptTokensDetails.CachedReadTokens != 2 { + t.Fatalf("precached prompt tokens were not mapped: %#v", response.Usage.PromptTokensDetails) + } + if len(response.Choices) != 1 || response.Choices[0].ChatNonStreamResponseChoice == nil { + t.Fatalf("unexpected choices: %#v", response.Choices) + } + content := response.Choices[0].ChatNonStreamResponseChoice.Message.Content + if content == nil || content.ContentStr == nil || *content.ContentStr != "Здравствуйте" { + t.Fatalf("content mismatch: %#v", content) + } + assistant := response.Choices[0].ChatNonStreamResponseChoice.Message.ChatAssistantMessage + if assistant == nil || assistant.Reasoning == nil || *assistant.Reasoning != "Думаю" { + t.Fatalf("reasoning_content was not mapped: %#v", assistant) + } + if len(assistant.ReasoningDetails) != 1 || assistant.ReasoningDetails[0].Text == nil || *assistant.ReasoningDetails[0].Text != "Думаю" { + t.Fatalf("reasoning details mismatch: %#v", assistant.ReasoningDetails) + } + if got := ctx.Value(schemas.BifrostContextKeyProviderResponseHeaders); got == nil { + t.Fatal("provider response headers were not stored in context") + } +} + +func testGigaChatChatCompletionExecutesWithMTLSClientCertificate(t *testing.T) { + t.Parallel() + + var oauthRequests atomic.Int32 + var chatRequests atomic.Int32 + server, caBundleFile, certFile, keyFile := newGigaChatMTLSServer(t, http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) { + switch request.URL.Path { + case "/oauth", "/api/v2/oauth", "/v1/token", "/api/v1/token": + oauthRequests.Add(1) + t.Fatalf("unexpected token endpoint request: %s", request.URL.Path) + case "/v1/chat/completions": + chatRequests.Add(1) + if got := request.Header.Get("Authorization"); got != "" { + t.Fatalf("chat authorization header mismatch: got %q, want empty", got) + } + if request.TLS == nil || len(request.TLS.PeerCertificates) == 0 { + t.Fatal("expected client certificate on API request") + } + assertGigaChatChatRequestBody(t, request) + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{ + "id":"chatcmpl-mtls", + "choices":[{"index":0,"message":{"role":"assistant","content":"Здравствуйте"},"finish_reason":"stop"}], + "created":1700000000, + "model":"GigaChat", + "object":"chat.completion", + "usage":{"prompt_tokens":7,"completion_tokens":3,"total_tokens":10} + }`)) + default: + t.Fatalf("unexpected path: %s", request.URL.Path) + } + })) + + provider := newTestGigaChatChatProvider(t, server.URL) + ctx := testBifrostContext() + ctx.SetValue(schemas.BifrostContextKeyPassthroughExtraParams, true) + response, bifrostErr := provider.ChatCompletion(ctx, schemas.Key{ + GigaChatKeyConfig: &schemas.GigaChatKeyConfig{ + CertFile: certFile, + KeyFile: keyFile, + CABundleFile: caBundleFile, + }, + }, testGigaChatChatRequest()) + if bifrostErr != nil { + t.Fatalf("ChatCompletion returned error: %v", bifrostErr) + } + if oauthRequests.Load() != 0 { + t.Fatalf("oauth request count mismatch: got %d, want 0", oauthRequests.Load()) + } + if chatRequests.Load() != 1 { + t.Fatalf("chat request count mismatch: got %d, want 1", chatRequests.Load()) + } + if response.ID != "chatcmpl-mtls" { + t.Fatalf("response id mismatch: got %q", response.ID) + } +} + +func testGigaChatChatCompletionUploadsInlineImageAttachment(t *testing.T) { + t.Parallel() + + var uploadRequests atomic.Int32 + var chatRequests atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) { + switch request.URL.Path { + case "/v1/files": + uploadRequests.Add(1) + if got := request.Header.Get("Authorization"); got != "Bearer image-token" { + t.Fatalf("file upload authorization header mismatch: got %q", got) + } + if err := request.ParseMultipartForm(1024); err != nil { + t.Fatalf("failed to parse upload multipart form: %v", err) + } + if got := request.FormValue("purpose"); got != "general" { + t.Fatalf("upload purpose mismatch: got %q", got) + } + file, header, err := request.FormFile("file") + if err != nil { + t.Fatalf("failed to read uploaded file: %v", err) + } + defer file.Close() + fileBytes, err := io.ReadAll(file) + if err != nil { + t.Fatalf("failed to read uploaded bytes: %v", err) + } + if string(fileBytes) != "image-bytes" { + t.Fatalf("uploaded image bytes mismatch: %q", fileBytes) + } + if header.Filename != "image.jpg" { + t.Fatalf("uploaded image filename mismatch: got %q", header.Filename) + } + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"id":"uploaded-image","object":"file","bytes":11,"created_at":1700000000,"filename":"image.jpg","purpose":"general"}`)) + case "/v1/chat/completions": + chatRequests.Add(1) + body, err := io.ReadAll(request.Body) + if err != nil { + t.Fatalf("failed to read chat body: %v", err) + } + payload := assertGigaChatChatBodyAttachment(t, body, "uploaded-image") + bodyStr := string(body) + if strings.Contains(bodyStr, "data:image") || strings.Contains(bodyStr, "image_url") { + t.Fatalf("chat body leaked OpenAI image_url payload: %s", body) + } + if _, ok := payload["function_call"]; ok { + t.Fatalf("image-only attachment should not force function_call auto: %s", body) + } + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"choices":[{"index":0,"message":{"role":"assistant","content":"На изображении..."},"finish_reason":"stop"}],"model":"GigaChat","object":"chat.completion"}`)) + default: + t.Fatalf("unexpected path: %s", request.URL.Path) + } + })) + defer server.Close() + + prompt := "Что на изображении?" + imageURL := "data:image/jpg;base64,aW1hZ2UtYnl0ZXM=" + request := &schemas.BifrostChatRequest{ + Model: "GigaChat", + Input: []schemas.ChatMessage{ + { + Role: schemas.ChatMessageRoleUser, + Content: &schemas.ChatMessageContent{ + ContentBlocks: []schemas.ChatContentBlock{ + {Type: schemas.ChatContentBlockTypeText, Text: &prompt}, + {Type: schemas.ChatContentBlockTypeImage, ImageURLStruct: &schemas.ChatInputImage{URL: imageURL}}, + }, + }, + }, + }, + } + + provider := newTestGigaChatChatProvider(t, server.URL) + response, bifrostErr := provider.ChatCompletion(testBifrostContext(), testGigaChatAccessTokenKey("image-token"), request) + if bifrostErr != nil { + t.Fatalf("ChatCompletion returned error: %v", bifrostErr) + } + if response == nil { + t.Fatal("expected response, got nil") + } + if uploadRequests.Load() != 1 { + t.Fatalf("upload request count mismatch: got %d, want 1", uploadRequests.Load()) + } + if chatRequests.Load() != 1 { + t.Fatalf("chat request count mismatch: got %d, want 1", chatRequests.Load()) + } +} + +func testGigaChatChatCompletionUploadsInlineFileAttachment(t *testing.T) { + t.Parallel() + + var uploadRequests atomic.Int32 + var chatRequests atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) { + switch request.URL.Path { + case "/v1/files": + uploadRequests.Add(1) + if err := request.ParseMultipartForm(1024); err != nil { + t.Fatalf("failed to parse upload multipart form: %v", err) + } + file, header, err := request.FormFile("file") + if err != nil { + t.Fatalf("failed to read uploaded file: %v", err) + } + defer file.Close() + fileBytes, err := io.ReadAll(file) + if err != nil { + t.Fatalf("failed to read uploaded bytes: %v", err) + } + if string(fileBytes) != "%PDF test" { + t.Fatalf("uploaded file bytes mismatch: %q", fileBytes) + } + if header.Filename != "Day_2_v6.pdf" { + t.Fatalf("uploaded filename mismatch: got %q", header.Filename) + } + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"id":"uploaded-pdf","object":"file","bytes":9,"created_at":1700000000,"filename":"Day_2_v6.pdf","purpose":"general"}`)) + case "/v1/chat/completions": + chatRequests.Add(1) + body, err := io.ReadAll(request.Body) + if err != nil { + t.Fatalf("failed to read chat body: %v", err) + } + payload := assertGigaChatChatBodyAttachment(t, body, "uploaded-pdf") + bodyStr := string(body) + if got := payload["function_call"]; got != "auto" { + t.Fatalf("document attachment should enable function_call auto: got %#v body %s", got, body) + } + if strings.Contains(bodyStr, "file_data") || strings.Contains(bodyStr, "application/pdf;base64") { + t.Fatalf("chat body leaked OpenAI file payload: %s", body) + } + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"choices":[{"index":0,"message":{"role":"assistant","content":"Краткое содержание..."},"finish_reason":"stop"}],"model":"GigaChat","object":"chat.completion"}`)) + default: + t.Fatalf("unexpected path: %s", request.URL.Path) + } + })) + defer server.Close() + + prompt := "Create a comprehensive summary of this pdf" + filename := "Day_2_v6.pdf" + fileData := "data:application/pdf;base64,JVBERiB0ZXN0" + request := &schemas.BifrostChatRequest{ + Model: "GigaChat", + Input: []schemas.ChatMessage{ + { + Role: schemas.ChatMessageRoleUser, + Content: &schemas.ChatMessageContent{ + ContentBlocks: []schemas.ChatContentBlock{ + {Type: schemas.ChatContentBlockTypeText, Text: &prompt}, + { + Type: schemas.ChatContentBlockTypeFile, + File: &schemas.ChatInputFile{ + Filename: &filename, + FileData: &fileData, + }, + }, + }, + }, + }, + }, + } + + provider := newTestGigaChatChatProvider(t, server.URL) + response, bifrostErr := provider.ChatCompletion(testBifrostContext(), testGigaChatAccessTokenKey("file-token"), request) + if bifrostErr != nil { + t.Fatalf("ChatCompletion returned error: %v", bifrostErr) + } + if response == nil { + t.Fatal("expected response, got nil") + } + if uploadRequests.Load() != 1 { + t.Fatalf("upload request count mismatch: got %d, want 1", uploadRequests.Load()) + } + if chatRequests.Load() != 1 { + t.Fatalf("chat request count mismatch: got %d, want 1", chatRequests.Load()) + } +} + +func testGigaChatChatCompletionReusesUploadedAttachmentAfterBackendError(t *testing.T) { + t.Parallel() + + var uploadRequests atomic.Int32 + var chatRequests atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) { + switch request.URL.Path { + case "/v1/files": + uploadRequests.Add(1) + if err := request.ParseMultipartForm(1024); err != nil { + t.Fatalf("failed to parse upload multipart form: %v", err) + } + file, _, err := request.FormFile("file") + if err != nil { + t.Fatalf("failed to read uploaded file: %v", err) + } + defer file.Close() + fileBytes, err := io.ReadAll(file) + if err != nil { + t.Fatalf("failed to read uploaded bytes: %v", err) + } + if string(fileBytes) != "%PDF retry" { + t.Fatalf("uploaded file bytes mismatch: %q", fileBytes) + } + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"id":"uploaded-retry-pdf","object":"file","bytes":10,"created_at":1700000000,"filename":"retry.pdf","purpose":"general"}`)) + case "/v1/chat/completions": + requestIndex := chatRequests.Add(1) + body, err := io.ReadAll(request.Body) + if err != nil { + t.Fatalf("failed to read chat body: %v", err) + } + assertGigaChatChatBodyAttachment(t, body, "uploaded-retry-pdf") + if requestIndex == 1 { + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusInternalServerError) + _, _ = w.Write([]byte(`{"status":500,"message":"temporary backend failure"}`)) + return + } + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"choices":[{"index":0,"message":{"role":"assistant","content":"ok"},"finish_reason":"stop"}],"model":"GigaChat","object":"chat.completion"}`)) + default: + t.Fatalf("unexpected path: %s", request.URL.Path) + } + })) + defer server.Close() + + prompt := "Summarize this file." + filename := "retry.pdf" + fileData := "data:application/pdf;base64,JVBERiByZXRyeQ==" + request := &schemas.BifrostChatRequest{ + Model: "GigaChat", + Input: []schemas.ChatMessage{{ + Role: schemas.ChatMessageRoleUser, + Content: &schemas.ChatMessageContent{ + ContentBlocks: []schemas.ChatContentBlock{ + {Type: schemas.ChatContentBlockTypeText, Text: &prompt}, + { + Type: schemas.ChatContentBlockTypeFile, + File: &schemas.ChatInputFile{ + Filename: &filename, + FileData: &fileData, + }, + }, + }, + }, + }}, + } + + provider := newTestGigaChatChatProvider(t, server.URL) + ctx := testBifrostContext() + key := testGigaChatAccessTokenKey("file-token") + + firstResponse, firstErr := provider.ChatCompletion(ctx, key, request) + if firstResponse != nil { + t.Fatalf("expected nil response from first backend failure, got %#v", firstResponse) + } + if firstErr == nil || firstErr.StatusCode == nil || *firstErr.StatusCode != http.StatusInternalServerError { + t.Fatalf("expected first backend 500, got %#v", firstErr) + } + + response, bifrostErr := provider.ChatCompletion(ctx, key, request) + if bifrostErr != nil { + t.Fatalf("second ChatCompletion returned error: %v", bifrostErr) + } + if response == nil { + t.Fatal("expected second response, got nil") + } + if uploadRequests.Load() != 1 { + t.Fatalf("upload request count mismatch: got %d, want 1", uploadRequests.Load()) + } + if chatRequests.Load() != 2 { + t.Fatalf("chat request count mismatch: got %d, want 2", chatRequests.Load()) + } +} + +func testGigaChatChatCompletionReusesSuccessfulAttachmentAfterPartialUploadFailure(t *testing.T) { + t.Parallel() + + var firstFileUploads atomic.Int32 + var secondFileUploads atomic.Int32 + var chatRequests atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) { + switch request.URL.Path { + case "/v1/files": + if err := request.ParseMultipartForm(1024); err != nil { + t.Fatalf("failed to parse upload multipart form: %v", err) + } + file, header, err := request.FormFile("file") + if err != nil { + t.Fatalf("failed to read uploaded file: %v", err) + } + defer file.Close() + fileBytes, err := io.ReadAll(file) + if err != nil { + t.Fatalf("failed to read uploaded bytes: %v", err) + } + + w.Header().Set("Content-Type", "application/json") + switch header.Filename { + case "first.txt": + firstFileUploads.Add(1) + if string(fileBytes) != "first attachment" { + t.Fatalf("first uploaded file bytes mismatch: %q", fileBytes) + } + _, _ = w.Write([]byte(`{"id":"uploaded-first","object":"file","bytes":16,"created_at":1700000000,"filename":"first.txt","purpose":"general"}`)) + case "second.txt": + uploadIndex := secondFileUploads.Add(1) + if string(fileBytes) != "second attachment" { + t.Fatalf("second uploaded file bytes mismatch: %q", fileBytes) + } + if uploadIndex == 1 { + w.WriteHeader(http.StatusInternalServerError) + _, _ = w.Write([]byte(`{"status":500,"message":"temporary upload failure"}`)) + return + } + _, _ = w.Write([]byte(`{"id":"uploaded-second","object":"file","bytes":17,"created_at":1700000000,"filename":"second.txt","purpose":"general"}`)) + default: + t.Fatalf("unexpected uploaded filename: %q", header.Filename) + } + case "/v1/chat/completions": + chatRequests.Add(1) + body, err := io.ReadAll(request.Body) + if err != nil { + t.Fatalf("failed to read chat body: %v", err) + } + var payload map[string]interface{} + if err := json.Unmarshal(body, &payload); err != nil { + t.Fatalf("failed to unmarshal chat body %s: %v", body, err) + } + messages, ok := payload["messages"].([]interface{}) + if !ok || len(messages) != 1 { + t.Fatalf("messages mismatch: %#v", payload["messages"]) + } + message, ok := messages[0].(map[string]interface{}) + if !ok { + t.Fatalf("message shape mismatch: %#v", messages[0]) + } + attachments, ok := message["attachments"].([]interface{}) + if !ok || len(attachments) != 2 || attachments[0] != "uploaded-first" || attachments[1] != "uploaded-second" { + t.Fatalf("attachments mismatch: %#v body %s", message["attachments"], body) + } + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"choices":[{"index":0,"message":{"role":"assistant","content":"ok"},"finish_reason":"stop"}],"model":"GigaChat","object":"chat.completion"}`)) + default: + t.Fatalf("unexpected path: %s", request.URL.Path) + } + })) + defer server.Close() + + prompt := "Summarize both files." + firstFilename := "first.txt" + firstFileData := "data:text/plain;base64,Zmlyc3QgYXR0YWNobWVudA==" + secondFilename := "second.txt" + secondFileData := "data:text/plain;base64,c2Vjb25kIGF0dGFjaG1lbnQ=" + request := &schemas.BifrostChatRequest{ + Model: "GigaChat", + Input: []schemas.ChatMessage{{ + Role: schemas.ChatMessageRoleUser, + Content: &schemas.ChatMessageContent{ + ContentBlocks: []schemas.ChatContentBlock{ + {Type: schemas.ChatContentBlockTypeText, Text: &prompt}, + { + Type: schemas.ChatContentBlockTypeFile, + File: &schemas.ChatInputFile{ + Filename: &firstFilename, + FileData: &firstFileData, + }, + }, + { + Type: schemas.ChatContentBlockTypeFile, + File: &schemas.ChatInputFile{ + Filename: &secondFilename, + FileData: &secondFileData, + }, + }, + }, + }, + }}, + } + + provider := newTestGigaChatChatProvider(t, server.URL) + parentCtx, cancel := context.WithCancel(context.Background()) + defer cancel() + ctx := schemas.NewBifrostContext(parentCtx, schemas.NoDeadline) + key := testGigaChatAccessTokenKey("file-token") + + firstResponse, firstErr := provider.ChatCompletion(ctx, key, request) + if firstResponse != nil { + t.Fatalf("expected nil response from partial upload failure, got %#v", firstResponse) + } + if firstErr == nil || firstErr.StatusCode == nil || *firstErr.StatusCode != http.StatusInternalServerError { + t.Fatalf("expected partial upload 500, got %#v", firstErr) + } + + response, bifrostErr := provider.ChatCompletion(ctx, key, request) + if bifrostErr != nil { + t.Fatalf("second ChatCompletion returned error: %v", bifrostErr) + } + if response == nil { + t.Fatal("expected second response, got nil") + } + if firstFileUploads.Load() != 1 { + t.Fatalf("first attachment upload count mismatch: got %d, want 1", firstFileUploads.Load()) + } + if secondFileUploads.Load() != 2 { + t.Fatalf("second attachment upload count mismatch: got %d, want 2", secondFileUploads.Load()) + } + if chatRequests.Load() != 1 { + t.Fatalf("chat request count mismatch: got %d, want 1", chatRequests.Load()) + } +} + +func testGigaChatChatCompletionDoesNotReuseUploadedAttachmentAcrossIndependentRequests(t *testing.T) { + t.Parallel() + + var uploadRequests atomic.Int32 + var chatRequests atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) { + switch request.URL.Path { + case "/v1/files": + uploadIndex := uploadRequests.Add(1) + if err := request.ParseMultipartForm(1024); err != nil { + t.Fatalf("failed to parse upload multipart form: %v", err) + } + file, _, err := request.FormFile("file") + if err != nil { + t.Fatalf("failed to read uploaded file: %v", err) + } + defer file.Close() + fileBytes, err := io.ReadAll(file) + if err != nil { + t.Fatalf("failed to read uploaded bytes: %v", err) + } + if string(fileBytes) != "%PDF independent" { + t.Fatalf("uploaded file bytes mismatch: %q", fileBytes) + } + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"id":"uploaded-independent-` + formatInt32(uploadIndex) + `","object":"file","bytes":16,"created_at":1700000000,"filename":"independent.pdf","purpose":"general"}`)) + case "/v1/chat/completions": + chatIndex := chatRequests.Add(1) + body, err := io.ReadAll(request.Body) + if err != nil { + t.Fatalf("failed to read chat body: %v", err) + } + assertGigaChatChatBodyAttachment(t, body, "uploaded-independent-"+formatInt32(chatIndex)) + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"choices":[{"index":0,"message":{"role":"assistant","content":"ok"},"finish_reason":"stop"}],"model":"GigaChat","object":"chat.completion"}`)) + default: + t.Fatalf("unexpected path: %s", request.URL.Path) + } + })) + defer server.Close() + + provider := newTestGigaChatChatProvider(t, server.URL) + ctx := testBifrostContext() + key := testGigaChatAccessTokenKey("file-token") + + firstResponse, firstErr := provider.ChatCompletion(ctx, key, testGigaChatInlineFileChatRequest("independent.pdf", "data:application/pdf;base64,JVBERiBpbmRlcGVuZGVudA==")) + if firstErr != nil { + t.Fatalf("first ChatCompletion returned error: %v", firstErr) + } + if firstResponse == nil { + t.Fatal("expected first response, got nil") + } + + secondResponse, secondErr := provider.ChatCompletion(ctx, key, testGigaChatInlineFileChatRequest("independent.pdf", "data:application/pdf;base64,JVBERiBpbmRlcGVuZGVudA==")) + if secondErr != nil { + t.Fatalf("second ChatCompletion returned error: %v", secondErr) + } + if secondResponse == nil { + t.Fatal("expected second response, got nil") + } + + if uploadRequests.Load() != 2 { + t.Fatalf("upload request count mismatch: got %d, want 2", uploadRequests.Load()) + } + if chatRequests.Load() != 2 { + t.Fatalf("chat request count mismatch: got %d, want 2", chatRequests.Load()) + } +} + +func testGigaChatChatCompletionDoesNotCacheFailedAttachmentUpload(t *testing.T) { + t.Parallel() + + var uploadRequests atomic.Int32 + var chatRequests atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) { + switch request.URL.Path { + case "/v1/files": + uploadIndex := uploadRequests.Add(1) + if err := request.ParseMultipartForm(1024); err != nil { + t.Fatalf("failed to parse upload multipart form: %v", err) + } + file, _, err := request.FormFile("file") + if err != nil { + t.Fatalf("failed to read uploaded file: %v", err) + } + defer file.Close() + fileBytes, err := io.ReadAll(file) + if err != nil { + t.Fatalf("failed to read uploaded bytes: %v", err) + } + if string(fileBytes) != "stale-secret-inline" { + t.Fatalf("uploaded file bytes mismatch: %q", fileBytes) + } + w.Header().Set("Content-Type", "application/json") + if uploadIndex == 1 { + w.WriteHeader(http.StatusInternalServerError) + _, _ = w.Write([]byte(`{"status":500,"message":"upload failed","id":"stale-file-id"}`)) + return + } + _, _ = w.Write([]byte(`{"id":"uploaded-after-error","object":"file","bytes":19,"created_at":1700000000,"filename":"secret.txt","purpose":"general"}`)) + case "/v1/chat/completions": + chatRequests.Add(1) + body, err := io.ReadAll(request.Body) + if err != nil { + t.Fatalf("failed to read chat body: %v", err) + } + assertGigaChatChatBodyAttachment(t, body, "uploaded-after-error") + if strings.Contains(string(body), "stale-file-id") { + t.Fatalf("chat body used stale file id: %s", body) + } + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"choices":[{"index":0,"message":{"role":"assistant","content":"ok"},"finish_reason":"stop"}],"model":"GigaChat","object":"chat.completion"}`)) + default: + t.Fatalf("unexpected path: %s", request.URL.Path) + } + })) + defer server.Close() + + provider := newTestGigaChatChatProvider(t, server.URL) + provider.sendBackRawRequest = true + provider.sendBackRawResponse = true + ctx := testBifrostContext() + key := testGigaChatAccessTokenKey("file-token") + request := testGigaChatInlineFileChatRequest("secret.txt", "data:text/plain;base64,c3RhbGUtc2VjcmV0LWlubGluZQ==") + + firstResponse, firstErr := provider.ChatCompletion(ctx, key, request) + if firstResponse != nil { + t.Fatalf("expected nil response from failed upload, got %#v", firstResponse) + } + if firstErr == nil || firstErr.StatusCode == nil || *firstErr.StatusCode != http.StatusInternalServerError { + t.Fatalf("expected upload 500, got %#v", firstErr) + } + if cacheID := ctx.Value(gigaChatAttachmentCacheKey); cacheID != nil { + t.Fatalf("failed attachment upload allocated context cache state: %#v", cacheID) + } + provider.attachmentCache.mu.Lock() + cacheEntries := len(provider.attachmentCache.entries) + provider.attachmentCache.mu.Unlock() + if cacheEntries != 0 { + t.Fatalf("failed attachment upload retained %d cache entries", cacheEntries) + } + firstErrorOutput := firstErr.String() + stringifyGigaChatRaw(firstErr.ExtraFields.RawRequest) + stringifyGigaChatRaw(firstErr.ExtraFields.RawResponse) + if strings.Contains(firstErrorOutput, "stale-secret-inline") || strings.Contains(firstErrorOutput, "c3RhbGUtc2VjcmV0LWlubGluZQ") { + t.Fatalf("failed upload leaked inline payload in error output: %s", firstErrorOutput) + } + + response, bifrostErr := provider.ChatCompletion(ctx, key, request) + if bifrostErr != nil { + t.Fatalf("second ChatCompletion returned error: %v", bifrostErr) + } + if response == nil { + t.Fatal("expected second response, got nil") + } + if uploadRequests.Load() != 2 { + t.Fatalf("upload request count mismatch: got %d, want 2", uploadRequests.Load()) + } + if chatRequests.Load() != 1 { + t.Fatalf("chat request count mismatch: got %d, want 1", chatRequests.Load()) + } +} + +func testGigaChatChatCompletionRejectsUnsupportedTools(t *testing.T) { + t.Parallel() + + text := "hello" + request := &schemas.BifrostChatRequest{ + Model: "GigaChat", + Input: []schemas.ChatMessage{ + {Role: schemas.ChatMessageRoleUser, Content: &schemas.ChatMessageContent{ContentStr: &text}}, + }, + Params: &schemas.ChatParameters{ + Tools: []schemas.ChatTool{ + { + Type: schemas.ChatToolTypeFunction, + Function: &schemas.ChatToolFunction{ + Name: "get_weather", + }, + }, + }, + }, + } + + _, err := ToGigaChatChatRequest(testBifrostContext(), request) + if err == nil { + t.Fatal("expected unsupported tools error, got nil") + } + if !strings.Contains(err.Error(), "tools") { + t.Fatalf("unexpected error: %v", err) + } +} + +func testGigaChatChatCompletionRejectsUnsupportedResponseFormat(t *testing.T) { + t.Parallel() + + responseFormat := interface{}(map[string]interface{}{"type": "json_object"}) + request := testGigaChatChatRequest() + request.Params.ResponseFormat = &responseFormat + + _, err := ToGigaChatChatRequest(testBifrostContext(), request) + if err == nil { + t.Fatal("expected unsupported response_format error, got nil") + } + if !strings.Contains(err.Error(), `response_format type "json_object" is not supported`) { + t.Fatalf("unexpected error: %v", err) + } +} + +func testGigaChatChatCompletionMapsProviderErrors(t *testing.T) { + t.Parallel() + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusBadRequest) + _, _ = w.Write([]byte(`{"status":400,"code":123,"message":"bad request"}`)) + })) + defer server.Close() + + provider := newTestGigaChatChatProvider(t, server.URL) + response, bifrostErr := provider.ChatCompletion(testBifrostContext(), testGigaChatAccessTokenKey("provider-error-token"), testGigaChatChatRequest()) + if response != nil { + t.Fatalf("expected nil response, got %#v", response) + } + if bifrostErr == nil { + t.Fatal("expected provider error, got nil") + } + if bifrostErr.StatusCode == nil || *bifrostErr.StatusCode != http.StatusBadRequest { + t.Fatalf("status mismatch: %#v", bifrostErr.StatusCode) + } + if bifrostErr.Error == nil || bifrostErr.Error.Message != "bad request" { + t.Fatalf("message mismatch: %#v", bifrostErr.Error) + } + if bifrostErr.Error.Code == nil || *bifrostErr.Error.Code != "123" { + t.Fatalf("code mismatch: %#v", bifrostErr.Error) + } + assertNoGigaChatSecretLeak(t, bifrostErr.String()) +} + +func testGigaChatChatCompletionRefreshesTokenAfterUnauthorized(t *testing.T) { + t.Parallel() + + var tokenRequests atomic.Int32 + var chatRequests atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) { + switch request.URL.Path { + case "/oauth": + tokenIndex := tokenRequests.Add(1) + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(fmt.Sprintf(`{"access_token":"token-%d","expires_at":1893456000}`, tokenIndex))) + case "/v1/chat/completions": + chatIndex := chatRequests.Add(1) + wantAuthorization := fmt.Sprintf("Bearer token-%d", chatIndex) + if got := request.Header.Get("Authorization"); got != wantAuthorization { + t.Fatalf("authorization header mismatch on request %d: got %q, want %q", chatIndex, got, wantAuthorization) + } + w.Header().Set("Content-Type", "application/json") + if chatIndex == 1 { + w.WriteHeader(http.StatusUnauthorized) + _, _ = w.Write([]byte(`{"status":401,"message":"expired token"}`)) + return + } + _, _ = w.Write([]byte(`{"choices":[{"index":0,"message":{"role":"assistant","content":"ok"},"finish_reason":"stop"}],"model":"GigaChat","object":"chat.completion"}`)) + default: + t.Fatalf("unexpected path: %s", request.URL.Path) + } + })) + defer server.Close() + + provider := newTestGigaChatChatProvider(t, server.URL) + response, bifrostErr := provider.ChatCompletion(testBifrostContext(), testGigaChatOAuthKey(server.URL+"/oauth", "", "test-credentials"), testGigaChatChatRequest()) + if bifrostErr != nil { + t.Fatalf("ChatCompletion returned error: %v", bifrostErr) + } + if response == nil { + t.Fatal("expected response, got nil") + } + if tokenRequests.Load() != 2 { + t.Fatalf("token request count mismatch: got %d, want 2", tokenRequests.Load()) + } + if chatRequests.Load() != 2 { + t.Fatalf("chat request count mismatch: got %d, want 2", chatRequests.Load()) + } +} + +func testGigaChatChatCompletionDoesNotDoubleExchangeExpiredTokenOnRefresh(t *testing.T) { + t.Parallel() + + var nowUnix atomic.Int64 + nowUnix.Store(time.Unix(1_700_000_000, 0).Unix()) + currentNow := func() time.Time { + return time.Unix(nowUnix.Load(), 0) + } + var tokenRequests atomic.Int32 + var chatRequests atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) { + switch request.URL.Path { + case "/oauth": + tokenIndex := tokenRequests.Add(1) + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(fmt.Sprintf(`{"access_token":"token-%d","expires_at":%d}`, tokenIndex, currentNow().Add(30*time.Minute).Unix()))) + case "/v1/chat/completions": + chatIndex := chatRequests.Add(1) + wantAuthorization := fmt.Sprintf("Bearer token-%d", chatIndex) + if got := request.Header.Get("Authorization"); got != wantAuthorization { + t.Fatalf("authorization header mismatch on request %d: got %q, want %q", chatIndex, got, wantAuthorization) + } + w.Header().Set("Content-Type", "application/json") + if chatIndex == 1 { + nowUnix.Add(int64(31 * time.Minute / time.Second)) + w.WriteHeader(http.StatusUnauthorized) + _, _ = w.Write([]byte(`{"status":401,"message":"expired token"}`)) + return + } + _, _ = w.Write([]byte(`{"choices":[{"index":0,"message":{"role":"assistant","content":"ok"},"finish_reason":"stop"}],"model":"GigaChat","object":"chat.completion"}`)) + default: + t.Fatalf("unexpected path: %s", request.URL.Path) + } + })) + defer server.Close() + + provider := newTestGigaChatChatProvider(t, server.URL) + provider.tokenCache = newGigaChatTokenCache(currentNow) + + response, bifrostErr := provider.ChatCompletion(testBifrostContext(), testGigaChatOAuthKey(server.URL+"/oauth", "", "test-credentials"), testGigaChatChatRequest()) + if bifrostErr != nil { + t.Fatalf("ChatCompletion returned error: %v", bifrostErr) + } + if response == nil || len(response.Choices) != 1 { + t.Fatalf("unexpected response: %#v", response) + } + if tokenRequests.Load() != 2 { + t.Fatalf("token request count mismatch: got %d, want 2", tokenRequests.Load()) + } + if chatRequests.Load() != 2 { + t.Fatalf("chat request count mismatch: got %d, want 2", chatRequests.Load()) + } +} + +func testGigaChatChatCompletionStreamsSSEChunks(t *testing.T) { + t.Parallel() + + var tokenRequests atomic.Int32 + var streamRequests atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) { + switch request.URL.Path { + case "/oauth": + tokenRequests.Add(1) + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"access_token":"stream-access-token","expires_at":1893456000}`)) + case "/v1/chat/completions": + streamRequests.Add(1) + if got := request.Header.Get("Authorization"); got != "Bearer stream-access-token" { + t.Fatalf("stream authorization header mismatch: got %q", got) + } + if strings.Contains(request.Header.Get("Authorization"), "super-secret-credentials") { + t.Fatal("stream request leaked OAuth credentials") + } + assertGigaChatChatStreamRequestBody(t, request) + w.Header().Set("Content-Type", "text/event-stream") + w.Header().Set("X-Request-ID", "stream-request-id") + _, _ = w.Write([]byte("data: {\"id\":\"chatcmpl-stream\",\"choices\":[{\"index\":0,\"delta\":{\"role\":\"assistant\",\"reasoning_content\":\"Думаю\",\"content\":\"З\"}}],\"created\":1700000000,\"model\":\"GigaChat\",\"object\":\"chat.completion\"}\n\n")) + _, _ = w.Write([]byte("data: {\"id\":\"chatcmpl-stream\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"дравствуйте\"}}],\"created\":1700000000,\"model\":\"GigaChat\",\"object\":\"chat.completion\"}\n\n")) + _, _ = w.Write([]byte("data: {\"id\":\"chatcmpl-stream\",\"choices\":[{\"index\":0,\"delta\":{},\"finish_reason\":\"stop\"}],\"created\":1700000000,\"model\":\"GigaChat\",\"object\":\"chat.completion\",\"usage\":{\"prompt_tokens\":7,\"completion_tokens\":3,\"total_tokens\":10}}\n\n")) + _, _ = w.Write([]byte("data: [DONE]\n\n")) + default: + t.Fatalf("unexpected path: %s", request.URL.Path) + } + })) + defer server.Close() + + provider := newTestGigaChatChatProvider(t, server.URL) + provider.sendBackRawRequest = true + provider.sendBackRawResponse = true + ctx := testBifrostContext() + ctx.SetValue(schemas.BifrostContextKeyCaptureRawRequest, true) + ctx.SetValue(schemas.BifrostContextKeyCaptureRawResponse, true) + ctx.SetValue(schemas.BifrostContextKeyPassthroughExtraParams, true) + + stream, bifrostErr := provider.ChatCompletionStream(ctx, testGigaChatPostHookRunner, nil, testGigaChatOAuthKey(server.URL+"/oauth", "", "super-secret-credentials"), testGigaChatChatRequest()) + if bifrostErr != nil { + t.Fatalf("ChatCompletionStream returned error: %v", bifrostErr) + } + + chunks := collectGigaChatStreamChunks(t, stream) + if tokenRequests.Load() != 1 { + t.Fatalf("token request count mismatch: got %d, want 1", tokenRequests.Load()) + } + if streamRequests.Load() != 1 { + t.Fatalf("stream request count mismatch: got %d, want 1", streamRequests.Load()) + } + if len(chunks) != 3 { + t.Fatalf("chunk count mismatch: got %d, want 3: %#v", len(chunks), chunks) + } + + assertGigaChatStreamContentChunk(t, chunks[0], "З") + assertGigaChatStreamReasoningChunk(t, chunks[0], "Думаю") + if chunks[0].BifrostChatResponse.ExtraFields.RawResponse == nil { + t.Fatal("expected raw response on content chunk") + } + assertGigaChatStreamContentChunk(t, chunks[1], "дравствуйте") + if chunks[1].BifrostChatResponse.ExtraFields.RawResponse == nil { + t.Fatal("expected raw response on content chunk") + } + finalChunk := chunks[2].BifrostChatResponse + if finalChunk == nil { + t.Fatalf("final chunk missing chat response: %#v", chunks[2]) + } + if finalChunk.ExtraFields.Provider != schemas.GigaChat { + t.Fatalf("provider mismatch: got %q, want %q", finalChunk.ExtraFields.Provider, schemas.GigaChat) + } + if finalChunk.Usage == nil || finalChunk.Usage.TotalTokens != 10 { + t.Fatalf("usage mismatch: %#v", finalChunk.Usage) + } + if len(finalChunk.Choices) != 1 || finalChunk.Choices[0].FinishReason == nil || *finalChunk.Choices[0].FinishReason != "stop" { + t.Fatalf("finish reason mismatch: %#v", finalChunk.Choices) + } + if finalChunk.ExtraFields.RawRequest == nil { + t.Fatalf("expected raw request on final chunk, got request=%#v", finalChunk.ExtraFields.RawRequest) + } + if got := ctx.Value(schemas.BifrostContextKeyProviderResponseHeaders); got == nil { + t.Fatal("provider response headers were not stored in context") + } +} + +func testGigaChatChatCompletionStreamFinalizesLargeResponsePassthrough(t *testing.T) { + t.Parallel() + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) { + if request.URL.Path != "/v1/chat/completions" { + t.Fatalf("unexpected path: %s", request.URL.Path) + } + assertGigaChatChatStreamRequestBody(t, request) + w.Header().Set("Content-Type", "text/event-stream") + _, _ = w.Write([]byte("data: {\"id\":\"chatcmpl-large\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"large\"}}],\"created\":1700000000,\"model\":\"GigaChat\",\"object\":\"chat.completion\"}\n\n")) + _, _ = w.Write([]byte("data: [DONE]\n\n")) + })) + defer server.Close() + + provider := newTestGigaChatChatProvider(t, server.URL) + ctx := testBifrostContext() + ctx.SetValue(schemas.BifrostContextKeyLargePayloadMode, true) + ctx.SetValue(schemas.BifrostContextKeyPassthroughExtraParams, true) + + var finalizerCalls atomic.Int32 + stream, bifrostErr := provider.ChatCompletionStream(ctx, testGigaChatPostHookRunner, func(context.Context) { + finalizerCalls.Add(1) + }, testGigaChatAccessTokenKey("chat-stream-token"), testGigaChatChatRequest()) + if bifrostErr != nil { + t.Fatalf("ChatCompletionStream returned error: %v", bifrostErr) + } + + select { + case chunk, ok := <-stream: + if ok { + t.Fatalf("passthrough stream channel should be closed without chunks, got %#v", chunk) + } + case <-time.After(time.Second): + t.Fatal("timed out waiting for passthrough stream channel to close") + } + if got := finalizerCalls.Load(); got != 0 { + t.Fatalf("finalizer ran before passthrough delivery: got %d calls", got) + } + + reader, ok := ctx.Value(schemas.BifrostContextKeyLargeResponseReader).(io.ReadCloser) + if !ok || reader == nil { + t.Fatalf("large response reader missing from context: %#v", ctx.Value(schemas.BifrostContextKeyLargeResponseReader)) + } + finalizingReader, ok := reader.(*gigaChatPassthroughReadCloser) + if !ok { + t.Fatalf("passthrough reader type mismatch: %T", reader) + } + largeReader, ok := finalizingReader.ReadCloser.(*providerUtils.LargeResponseReader) + if !ok { + t.Fatalf("wrapped large response reader type mismatch: %T", finalizingReader.ReadCloser) + } + + body, err := io.ReadAll(reader) + if err != nil { + t.Fatalf("failed to read passthrough body: %v", err) + } + if !strings.Contains(string(body), `"id":"chatcmpl-large"`) { + t.Fatalf("passthrough body mismatch: %s", body) + } + if got := finalizerCalls.Load(); got != 0 { + t.Fatalf("finalizer ran before passthrough reader close: got %d calls", got) + } + if err := reader.Close(); err != nil { + t.Fatalf("failed to close large response reader: %v", err) + } + if got := finalizerCalls.Load(); got != 1 { + t.Fatalf("finalizer calls mismatch after passthrough delivery: got %d, want 1", got) + } + if ended, _ := ctx.Value(schemas.BifrostContextKeyStreamEndIndicator).(bool); !ended { + t.Fatal("passthrough stream was not marked complete before finalization") + } + if largeReader.Resp != nil { + t.Fatal("large response reader did not release its fasthttp response") + } +} + +func testGigaChatChatCompletionMapsStreamingProviderErrors(t *testing.T) { + t.Parallel() + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusBadRequest) + _, _ = w.Write([]byte(`{"status":400,"code":123,"message":"bad stream request"}`)) + })) + defer server.Close() + + provider := newTestGigaChatChatProvider(t, server.URL) + stream, bifrostErr := provider.ChatCompletionStream(testBifrostContext(), testGigaChatPostHookRunner, nil, testGigaChatAccessTokenKey("provider-error-token"), testGigaChatChatRequest()) + if stream != nil { + t.Fatalf("expected nil stream, got %#v", stream) + } + if bifrostErr == nil { + t.Fatal("expected provider error, got nil") + } + if bifrostErr.StatusCode == nil || *bifrostErr.StatusCode != http.StatusBadRequest { + t.Fatalf("status mismatch: %#v", bifrostErr.StatusCode) + } + if bifrostErr.Error == nil || bifrostErr.Error.Message != "bad stream request" { + t.Fatalf("message mismatch: %#v", bifrostErr.Error) + } + assertNoGigaChatSecretLeak(t, bifrostErr.String()) +} + +func testGigaChatChatCompletionMapsStreamingErrorEvents(t *testing.T) { + t.Parallel() + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Type", "text/event-stream") + _, _ = w.Write([]byte("data: {\"status\":429,\"code\":42901,\"message\":\"rate limit\"}\n\n")) + _, _ = w.Write([]byte("data: [DONE]\n\n")) + })) + defer server.Close() + + provider := newTestGigaChatChatProvider(t, server.URL) + stream, bifrostErr := provider.ChatCompletionStream(testBifrostContext(), testGigaChatPostHookRunner, nil, testGigaChatAccessTokenKey("provider-error-token"), testGigaChatChatRequest()) + if bifrostErr != nil { + t.Fatalf("ChatCompletionStream returned error before stream: %v", bifrostErr) + } + chunks := collectGigaChatStreamChunks(t, stream) + if len(chunks) != 1 || chunks[0].BifrostError == nil { + t.Fatalf("expected one error chunk, got %#v", chunks) + } + streamErr := chunks[0].BifrostError + if streamErr.StatusCode == nil || *streamErr.StatusCode != http.StatusTooManyRequests { + t.Fatalf("status mismatch: %#v", streamErr.StatusCode) + } + if streamErr.Error == nil || streamErr.Error.Message != "rate limit" { + t.Fatalf("message mismatch: %#v", streamErr.Error) + } + if streamErr.Error.Code == nil || *streamErr.Error.Code != "42901" { + t.Fatalf("code mismatch: %#v", streamErr.Error) + } +} + +func testGigaChatChatCompletionStreamRefreshesTokenAfterUnauthorized(t *testing.T) { + t.Parallel() + + var tokenRequests atomic.Int32 + var streamRequests atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) { + switch request.URL.Path { + case "/oauth": + tokenIndex := tokenRequests.Add(1) + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(fmt.Sprintf(`{"access_token":"stream-token-%d","expires_at":1893456000}`, tokenIndex))) + case "/v1/chat/completions": + streamIndex := streamRequests.Add(1) + wantAuthorization := fmt.Sprintf("Bearer stream-token-%d", streamIndex) + if got := request.Header.Get("Authorization"); got != wantAuthorization { + t.Fatalf("authorization header mismatch on stream request %d: got %q, want %q", streamIndex, got, wantAuthorization) + } + if streamIndex == 1 { + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusUnauthorized) + _, _ = w.Write([]byte(`{"status":401,"message":"expired token"}`)) + return + } + w.Header().Set("Content-Type", "text/event-stream") + _, _ = w.Write([]byte("data: {\"choices\":[{\"index\":0,\"delta\":{\"content\":\"ok\"}}],\"model\":\"GigaChat\",\"object\":\"chat.completion\"}\n\n")) + _, _ = w.Write([]byte("data: [DONE]\n\n")) + default: + t.Fatalf("unexpected path: %s", request.URL.Path) + } + })) + defer server.Close() + + provider := newTestGigaChatChatProvider(t, server.URL) + stream, bifrostErr := provider.ChatCompletionStream(testBifrostContext(), testGigaChatPostHookRunner, nil, testGigaChatOAuthKey(server.URL+"/oauth", "", "test-credentials"), testGigaChatChatRequest()) + if bifrostErr != nil { + t.Fatalf("ChatCompletionStream returned error: %v", bifrostErr) + } + chunks := collectGigaChatStreamChunks(t, stream) + if tokenRequests.Load() != 2 { + t.Fatalf("token request count mismatch: got %d, want 2", tokenRequests.Load()) + } + if streamRequests.Load() != 2 { + t.Fatalf("stream request count mismatch: got %d, want 2", streamRequests.Load()) + } + if len(chunks) != 2 { + t.Fatalf("chunk count mismatch: got %d, want 2: %#v", len(chunks), chunks) + } + assertGigaChatStreamContentChunk(t, chunks[0], "ok") +} + +func testGigaChatChatCompletionStreamHandlesContextCancellation(t *testing.T) { + t.Parallel() + + firstChunkWritten := make(chan struct{}) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) { + if request.URL.Path != "/v1/chat/completions" { + t.Fatalf("unexpected path: %s", request.URL.Path) + } + w.Header().Set("Content-Type", "text/event-stream") + _, _ = w.Write([]byte("data: {\"choices\":[{\"index\":0,\"delta\":{\"content\":\"partial\"}}],\"model\":\"GigaChat\",\"object\":\"chat.completion\"}\n\n")) + if flusher, ok := w.(http.Flusher); ok { + flusher.Flush() + } + close(firstChunkWritten) + <-request.Context().Done() + })) + defer server.Close() + + provider := newTestGigaChatChatProvider(t, server.URL) + ctx, cancel := schemas.NewBifrostContextWithCancel(context.Background()) + stream, bifrostErr := provider.ChatCompletionStream(ctx, testGigaChatPostHookRunner, nil, testGigaChatAccessTokenKey("stream-token"), testGigaChatChatRequest()) + if bifrostErr != nil { + t.Fatalf("ChatCompletionStream returned error: %v", bifrostErr) + } + + select { + case <-firstChunkWritten: + case <-time.After(time.Second): + t.Fatal("timed out waiting for first stream chunk") + } + + firstChunk := <-stream + assertGigaChatStreamContentChunk(t, firstChunk, "partial") + cancel() + select { + case <-ctx.Done(): + case <-time.After(time.Second): + t.Fatal("timed out waiting for context cancellation") + } + + streamClosed := make(chan struct{}) + go func() { + for range stream { + } + close(streamClosed) + }() + + select { + case <-streamClosed: + case <-time.After(time.Second): + t.Fatal("timed out waiting for stream to close after context cancellation") + } +} + +func newTestGigaChatChatProvider(t *testing.T, baseURL string) *GigaChatProvider { + t.Helper() + + provider, err := NewGigaChatProvider(&schemas.ProviderConfig{ + NetworkConfig: schemas.NetworkConfig{ + BaseURL: baseURL, + }, + }, nil) + if err != nil { + t.Fatalf("NewGigaChatProvider returned error: %v", err) + } + dialer := &net.Dialer{} + provider.client.Dial = func(addr string) (net.Conn, error) { + return dialer.Dial("tcp", addr) + } + provider.client.DialTimeout = nil + provider.streamingClient.Dial = provider.client.Dial + provider.streamingClient.DialTimeout = nil + return provider +} + +func testGigaChatPostHookRunner(_ *schemas.BifrostContext, response *schemas.BifrostResponse, bifrostErr *schemas.BifrostError) (*schemas.BifrostResponse, *schemas.BifrostError) { + return response, bifrostErr +} + +func testGigaChatAccessTokenKey(accessToken string) schemas.Key { + return schemas.Key{ + GigaChatKeyConfig: &schemas.GigaChatKeyConfig{ + AccessToken: schemas.NewSecretVar(accessToken), + }, + } +} + +func testGigaChatChatRequest() *schemas.BifrostChatRequest { + maxTokens := 128 + text := "Привет" + return &schemas.BifrostChatRequest{ + Model: "GigaChat", + Input: []schemas.ChatMessage{ + { + Role: schemas.ChatMessageRoleUser, + Content: &schemas.ChatMessageContent{ContentStr: &text}, + }, + }, + Params: &schemas.ChatParameters{ + MaxCompletionTokens: &maxTokens, + ExtraParams: map[string]interface{}{ + "profanity_check": false, + }, + }, + } +} + +func testGigaChatInlineFileChatRequest(filename string, fileData string) *schemas.BifrostChatRequest { + prompt := "Summarize this file." + return &schemas.BifrostChatRequest{ + Model: "GigaChat", + Input: []schemas.ChatMessage{{ + Role: schemas.ChatMessageRoleUser, + Content: &schemas.ChatMessageContent{ + ContentBlocks: []schemas.ChatContentBlock{ + {Type: schemas.ChatContentBlockTypeText, Text: &prompt}, + { + Type: schemas.ChatContentBlockTypeFile, + File: &schemas.ChatInputFile{ + Filename: &filename, + FileData: &fileData, + }, + }, + }, + }, + }}, + } +} + +func assertGigaChatChatRequestBody(t *testing.T, request *http.Request) { + t.Helper() + + assertGigaChatChatRequestBodyWithStream(t, request, false) +} + +func assertGigaChatChatStreamRequestBody(t *testing.T, request *http.Request) { + t.Helper() + + assertGigaChatChatRequestBodyWithStream(t, request, true) +} + +func assertGigaChatChatRequestBodyWithStream(t *testing.T, request *http.Request, wantStream bool) { + t.Helper() + + if request.Method != http.MethodPost { + t.Fatalf("method mismatch: got %s, want POST", request.Method) + } + if got := request.Header.Get("Content-Type"); !strings.HasPrefix(got, "application/json") { + t.Fatalf("content type mismatch: got %q", got) + } + if got := request.Header.Get(gigaChatUserAgentHeader); got != gigaChatUserAgent { + t.Fatalf("user-agent mismatch: got %q, want %q", got, gigaChatUserAgent) + } + + body, err := io.ReadAll(request.Body) + if err != nil { + t.Fatalf("failed to read request body: %v", err) + } + var payload map[string]interface{} + if err := json.Unmarshal(body, &payload); err != nil { + t.Fatalf("failed to unmarshal request body %s: %v", body, err) + } + if got := payload["model"]; got != "GigaChat" { + t.Fatalf("model mismatch: got %#v", got) + } + if got := payload["max_tokens"]; got != float64(128) { + t.Fatalf("max_tokens mismatch: got %#v", got) + } + if _, ok := payload["max_completion_tokens"]; ok { + t.Fatalf("max_completion_tokens should not be sent: %s", body) + } + if got := payload["stream"]; got != wantStream { + t.Fatalf("stream mismatch: got %#v, want %v", got, wantStream) + } + if got := payload["profanity_check"]; got != false { + t.Fatalf("profanity_check mismatch: got %#v", got) + } + messages, ok := payload["messages"].([]interface{}) + if !ok || len(messages) != 1 { + t.Fatalf("messages mismatch: %#v", payload["messages"]) + } + message, ok := messages[0].(map[string]interface{}) + if !ok { + t.Fatalf("message shape mismatch: %#v", messages[0]) + } + if got := message["role"]; got != "user" { + t.Fatalf("message role mismatch: got %#v", got) + } + if got := message["content"]; got != "Привет" { + t.Fatalf("message content mismatch: got %#v", got) + } +} + +func assertGigaChatChatBodyAttachment(t *testing.T, body []byte, wantAttachment string) map[string]interface{} { + t.Helper() + + var payload map[string]interface{} + if err := json.Unmarshal(body, &payload); err != nil { + t.Fatalf("failed to unmarshal chat body %s: %v", body, err) + } + messages, ok := payload["messages"].([]interface{}) + if !ok || len(messages) != 1 { + t.Fatalf("messages mismatch: %#v", payload["messages"]) + } + message, ok := messages[0].(map[string]interface{}) + if !ok { + t.Fatalf("message shape mismatch: %#v", messages[0]) + } + attachments, ok := message["attachments"].([]interface{}) + if !ok || len(attachments) != 1 || attachments[0] != wantAttachment { + t.Fatalf("attachments mismatch: %#v body %s", message["attachments"], body) + } + return payload +} + +func assertGigaChatJSONSchemaResponseFormat(t *testing.T, responseFormat interface{}, wantTitle string, wantDescription string, wantStrict bool) { + t.Helper() + + formatMap, ok := schemas.SafeExtractOrderedMap(responseFormat) + if !ok || formatMap == nil { + t.Fatalf("response_format should be a JSON object: %#v", responseFormat) + } + if got, _ := formatMap.Get("type"); got != "json_schema" { + t.Fatalf("response_format type mismatch: got %#v", got) + } + if _, hasOpenAIWrapper := formatMap.Get("json_schema"); hasOpenAIWrapper { + t.Fatalf("response_format should not contain OpenAI json_schema wrapper: %#v", formatMap) + } + strictRaw, ok := formatMap.Get("strict") + if !ok { + t.Fatal("response_format strict is missing") + } + strict, ok := schemas.SafeExtractBool(strictRaw) + if !ok || strict != wantStrict { + t.Fatalf("response_format strict mismatch: got %#v, want %v", strictRaw, wantStrict) + } + schemaRaw, ok := formatMap.Get("schema") + if !ok { + t.Fatal("response_format schema is missing") + } + schemaMap, ok := schemas.SafeExtractOrderedMap(schemaRaw) + if !ok || schemaMap == nil { + t.Fatalf("response_format schema should be a JSON object: %#v", schemaRaw) + } + if got, _ := schemaMap.Get("type"); got != "object" { + t.Fatalf("schema type mismatch: got %#v", got) + } + if wantTitle != "" { + if got, _ := schemaMap.Get("title"); got != wantTitle { + t.Fatalf("schema title mismatch: got %#v, want %q", got, wantTitle) + } + } + if wantDescription != "" { + if got, _ := schemaMap.Get("description"); got != wantDescription { + t.Fatalf("schema description mismatch: got %#v, want %q", got, wantDescription) + } + } +} + +func collectGigaChatStreamChunks(t *testing.T, stream chan *schemas.BifrostStreamChunk) []*schemas.BifrostStreamChunk { + t.Helper() + + chunks := make([]*schemas.BifrostStreamChunk, 0) + for chunk := range stream { + chunks = append(chunks, chunk) + } + return chunks +} + +func assertGigaChatStreamContentChunk(t *testing.T, chunk *schemas.BifrostStreamChunk, wantContent string) { + t.Helper() + + if chunk == nil || chunk.BifrostChatResponse == nil { + t.Fatalf("missing chat stream response: %#v", chunk) + } + response := chunk.BifrostChatResponse + if response.ExtraFields.Provider != schemas.GigaChat { + t.Fatalf("provider mismatch: got %q, want %q", response.ExtraFields.Provider, schemas.GigaChat) + } + if len(response.Choices) != 1 || response.Choices[0].ChatStreamResponseChoice == nil || response.Choices[0].ChatStreamResponseChoice.Delta == nil { + t.Fatalf("unexpected choices: %#v", response.Choices) + } + content := response.Choices[0].ChatStreamResponseChoice.Delta.Content + if content == nil || *content != wantContent { + t.Fatalf("content mismatch: got %#v, want %q", content, wantContent) + } +} + +func assertGigaChatStreamReasoningChunk(t *testing.T, chunk *schemas.BifrostStreamChunk, wantReasoning string) { + t.Helper() + + if chunk == nil || chunk.BifrostChatResponse == nil { + t.Fatalf("missing chat stream response: %#v", chunk) + } + response := chunk.BifrostChatResponse + if len(response.Choices) != 1 || response.Choices[0].ChatStreamResponseChoice == nil || response.Choices[0].ChatStreamResponseChoice.Delta == nil { + t.Fatalf("unexpected choices: %#v", response.Choices) + } + delta := response.Choices[0].ChatStreamResponseChoice.Delta + if delta.Reasoning == nil || *delta.Reasoning != wantReasoning { + t.Fatalf("reasoning mismatch: got %#v, want %q", delta.Reasoning, wantReasoning) + } + if len(delta.ReasoningDetails) != 1 || delta.ReasoningDetails[0].Text == nil || *delta.ReasoningDetails[0].Text != wantReasoning { + t.Fatalf("reasoning details mismatch: %#v", delta.ReasoningDetails) + } +} diff --git a/core/providers/gigachat/count_tokens_test.go b/core/providers/gigachat/count_tokens_test.go new file mode 100644 index 00000000000..668118f8768 --- /dev/null +++ b/core/providers/gigachat/count_tokens_test.go @@ -0,0 +1,440 @@ +package gigachat + +import ( + "encoding/json" + "fmt" + "net" + "net/http" + "net/http/httptest" + "strings" + "sync/atomic" + "testing" + + "github.com/maximhq/bifrost/core/schemas" +) + +func TestGigaChatCountTokens(t *testing.T) { + t.Parallel() + + t.Run("ConverterMapsTextInput", testGigaChatCountTokensConverterMapsTextInput) + t.Run("ConverterRejectsEmptyText", testGigaChatCountTokensConverterRejectsEmptyText) + t.Run("ConverterRejectsImageContent", testGigaChatCountTokensConverterRejectsImageContent) + t.Run("ConverterRejectsFileContent", testGigaChatCountTokensConverterRejectsFileContent) + t.Run("ConverterRejectsAudioContent", testGigaChatCountTokensConverterRejectsAudioContent) + t.Run("ResponseMapsTokenSums", testGigaChatCountTokensResponseMapsTokenSums) + t.Run("ResponseAcceptsDataWrapper", testGigaChatCountTokensResponseAcceptsDataWrapper) + t.Run("ExecutesWithOAuthToken", testGigaChatCountTokensExecutesWithOAuthToken) + t.Run("MapsProviderErrors", testGigaChatCountTokensMapsProviderErrors) + t.Run("RefreshesTokenAfterUnauthorized", testGigaChatCountTokensRefreshesTokenAfterUnauthorized) +} + +func testGigaChatCountTokensConverterMapsTextInput(t *testing.T) { + t.Parallel() + + request := &schemas.BifrostResponsesRequest{ + Model: "GigaChat", + Input: []schemas.ResponsesMessage{ + { + Role: schemas.Ptr(schemas.ResponsesInputMessageRoleUser), + Content: &schemas.ResponsesMessageContent{ContentStr: schemas.Ptr("first")}, + }, + { + Role: schemas.Ptr(schemas.ResponsesInputMessageRoleUser), + Content: &schemas.ResponsesMessageContent{ContentBlocks: []schemas.ResponsesMessageContentBlock{ + { + Type: schemas.ResponsesInputMessageContentBlockTypeText, + Text: schemas.Ptr("second"), + }, + { + Type: schemas.ResponsesOutputMessageContentTypeText, + Text: schemas.Ptr("third"), + ResponsesOutputMessageContentText: &schemas.ResponsesOutputMessageContentText{ + Annotations: []schemas.ResponsesOutputMessageContentTextAnnotation{}, + LogProbs: []schemas.ResponsesOutputMessageContentTextLogProb{}, + }, + }, + { + ResponsesOutputMessageContentText: &schemas.ResponsesOutputMessageContentText{ + Annotations: []schemas.ResponsesOutputMessageContentTextAnnotation{}, + LogProbs: []schemas.ResponsesOutputMessageContentTextLogProb{}, + }, + }, + }}, + }, + { + Role: schemas.Ptr(schemas.ResponsesInputMessageRoleAssistant), + Content: &schemas.ResponsesMessageContent{ContentBlocks: []schemas.ResponsesMessageContentBlock{{ + Type: schemas.ResponsesOutputMessageContentTypeText, + Text: schemas.Ptr("assistant output"), + ResponsesOutputMessageContentText: &schemas.ResponsesOutputMessageContentText{ + Annotations: []schemas.ResponsesOutputMessageContentTextAnnotation{}, + LogProbs: []schemas.ResponsesOutputMessageContentTextLogProb{}, + }, + }}}, + }, + }, + } + + gigaChatReq, err := ToGigaChatCountTokensRequest(request) + if err != nil { + t.Fatalf("ToGigaChatCountTokensRequest returned error: %v", err) + } + if gigaChatReq.Model != "GigaChat" { + t.Fatalf("model mismatch: got %q", gigaChatReq.Model) + } + wantInput := []string{"first", "second", "third", "assistant output"} + if !equalStringSlices(gigaChatReq.Input, wantInput) { + t.Fatalf("input mismatch: got %#v, want %#v", gigaChatReq.Input, wantInput) + } + + body, err := json.Marshal(gigaChatReq) + if err != nil { + t.Fatalf("failed to marshal count tokens request: %v", err) + } + if !strings.Contains(string(body), `"model":"GigaChat"`) || !strings.Contains(string(body), `"input":["first","second","third","assistant output"]`) { + t.Fatalf("unexpected request body: %s", body) + } +} + +func testGigaChatCountTokensConverterRejectsEmptyText(t *testing.T) { + t.Parallel() + + _, err := ToGigaChatCountTokensRequest(&schemas.BifrostResponsesRequest{ + Model: "GigaChat", + Input: []schemas.ResponsesMessage{{ + Content: &schemas.ResponsesMessageContent{ContentStr: schemas.Ptr(" ")}, + }}, + }) + if err == nil { + t.Fatal("expected error, got nil") + } + if !strings.Contains(err.Error(), "empty") { + t.Fatalf("unexpected error: %v", err) + } +} + +func testGigaChatCountTokensConverterRejectsImageContent(t *testing.T) { + t.Parallel() + + _, err := ToGigaChatCountTokensRequest(&schemas.BifrostResponsesRequest{ + Model: "GigaChat", + Input: []schemas.ResponsesMessage{{ + Content: &schemas.ResponsesMessageContent{ContentBlocks: []schemas.ResponsesMessageContentBlock{{ + Type: schemas.ResponsesInputMessageContentBlockTypeImage, + ResponsesInputMessageContentBlockImage: &schemas.ResponsesInputMessageContentBlockImage{ + ImageURL: schemas.Ptr("https://example.com/image.png"), + }, + }}}, + }}, + }) + if err == nil { + t.Fatal("expected error, got nil") + } + if !strings.Contains(err.Error(), "file, image, or audio") { + t.Fatalf("unexpected error: %v", err) + } +} + +func testGigaChatCountTokensConverterRejectsFileContent(t *testing.T) { + t.Parallel() + + _, err := ToGigaChatCountTokensRequest(&schemas.BifrostResponsesRequest{ + Model: "GigaChat", + Input: []schemas.ResponsesMessage{{ + Content: &schemas.ResponsesMessageContent{ContentBlocks: []schemas.ResponsesMessageContentBlock{{ + Type: schemas.ResponsesInputMessageContentBlockTypeFile, + ResponsesInputMessageContentBlockFile: &schemas.ResponsesInputMessageContentBlockFile{ + FileData: schemas.Ptr("file-data"), + }, + }}}, + }}, + }) + if err == nil { + t.Fatal("expected error, got nil") + } + if !strings.Contains(err.Error(), "file, image, or audio") { + t.Fatalf("unexpected error: %v", err) + } +} + +func testGigaChatCountTokensConverterRejectsAudioContent(t *testing.T) { + t.Parallel() + + _, err := ToGigaChatCountTokensRequest(&schemas.BifrostResponsesRequest{ + Model: "GigaChat", + Input: []schemas.ResponsesMessage{{ + Content: &schemas.ResponsesMessageContent{ContentBlocks: []schemas.ResponsesMessageContentBlock{{ + Type: schemas.ResponsesInputMessageContentBlockTypeAudio, + Audio: &schemas.ResponsesInputMessageContentBlockAudio{ + Format: "mp3", + Data: "audio-data", + }, + }}}, + }}, + }) + if err == nil { + t.Fatal("expected error, got nil") + } + if !strings.Contains(err.Error(), "file, image, or audio") { + t.Fatalf("unexpected error: %v", err) + } +} + +func testGigaChatCountTokensResponseMapsTokenSums(t *testing.T) { + t.Parallel() + + response := ToBifrostCountTokensResponse(schemas.GigaChat, &GigaChatCountTokensResponse{ + Items: []GigaChatCountTokensItem{ + {Tokens: 3, Characters: 12}, + {Tokens: 5, Characters: 20}, + }, + }, "GigaChat") + if response == nil { + t.Fatal("response is nil") + } + if response.Object != "response.input_tokens" { + t.Fatalf("object mismatch: got %q", response.Object) + } + if response.Model != "GigaChat" { + t.Fatalf("model mismatch: got %q", response.Model) + } + if response.InputTokens != 8 { + t.Fatalf("input tokens mismatch: got %d, want 8", response.InputTokens) + } + if response.TotalTokens == nil || *response.TotalTokens != 8 { + t.Fatalf("total tokens mismatch: %#v", response.TotalTokens) + } + if !equalIntSlices(response.Tokens, []int{3, 5}) { + t.Fatalf("tokens mismatch: got %#v", response.Tokens) + } + if response.OutputTokens != nil { + t.Fatalf("expected nil output tokens, got %#v", response.OutputTokens) + } + if response.InputTokensDetails == nil || response.InputTokensDetails.TextTokens != 8 { + t.Fatalf("input token details mismatch: %#v", response.InputTokensDetails) + } + if response.ExtraFields.Provider != schemas.GigaChat { + t.Fatalf("provider mismatch: got %q", response.ExtraFields.Provider) + } +} + +func testGigaChatCountTokensResponseAcceptsDataWrapper(t *testing.T) { + t.Parallel() + + var gigaChatResponse GigaChatCountTokensResponse + if err := json.Unmarshal([]byte(`{"data":[{"tokens":2,"characters":7},{"tokens":4,"characters":13}]}`), &gigaChatResponse); err != nil { + t.Fatalf("failed to unmarshal data wrapper response: %v", err) + } + + response := ToBifrostCountTokensResponse(schemas.GigaChat, &gigaChatResponse, "GigaChat") + if response == nil || response.InputTokens != 6 || !equalIntSlices(response.Tokens, []int{2, 4}) { + t.Fatalf("unexpected response: %#v", response) + } +} + +func testGigaChatCountTokensExecutesWithOAuthToken(t *testing.T) { + t.Parallel() + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) { + switch request.URL.Path { + case "/oauth": + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"access_token":"count-tokens-access-token","expires_at":1893456000}`)) + case "/v1/tokens/count": + assertGigaChatCountTokensHTTPShape(t, request, "Bearer count-tokens-access-token") + assertGigaChatCountTokensRequestBody(t, request, []string{"first", "second"}) + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`[{"tokens":3,"characters":5},{"tokens":4,"characters":6}]`)) + default: + t.Fatalf("unexpected path: %s", request.URL.Path) + } + })) + defer server.Close() + + provider := newTestGigaChatCountTokensProvider(t, server.URL) + response, bifrostErr := provider.CountTokens(testBifrostContext(), testGigaChatOAuthKey(server.URL+"/oauth", "", "test-credentials"), testGigaChatCountTokensRequest()) + if bifrostErr != nil { + t.Fatalf("CountTokens returned error: %v", bifrostErr) + } + if response == nil { + t.Fatal("response is nil") + } + if response.InputTokens != 7 { + t.Fatalf("input tokens mismatch: got %d, want 7", response.InputTokens) + } + if response.TotalTokens == nil || *response.TotalTokens != 7 { + t.Fatalf("total tokens mismatch: %#v", response.TotalTokens) + } + if !equalIntSlices(response.Tokens, []int{3, 4}) { + t.Fatalf("tokens mismatch: got %#v", response.Tokens) + } + if response.ExtraFields.Provider != schemas.GigaChat { + t.Fatalf("provider mismatch: got %q", response.ExtraFields.Provider) + } +} + +func testGigaChatCountTokensMapsProviderErrors(t *testing.T) { + t.Parallel() + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) { + if request.URL.Path != "/v1/tokens/count" { + t.Fatalf("unexpected path: %s", request.URL.Path) + } + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusBadRequest) + _, _ = w.Write([]byte(`{"status":400,"code":"bad_count","message":"bad count tokens request"}`)) + })) + defer server.Close() + + provider := newTestGigaChatCountTokensProvider(t, server.URL) + response, bifrostErr := provider.CountTokens(testBifrostContext(), testGigaChatAccessTokenKey("count-tokens-error-token"), testGigaChatCountTokensRequest()) + if response != nil { + t.Fatalf("expected nil response, got %#v", response) + } + if bifrostErr == nil { + t.Fatal("expected provider error, got nil") + } + if bifrostErr.StatusCode == nil || *bifrostErr.StatusCode != http.StatusBadRequest { + t.Fatalf("status mismatch: %#v", bifrostErr.StatusCode) + } + if bifrostErr.Error == nil || bifrostErr.Error.Message != "bad count tokens request" { + t.Fatalf("message mismatch: %#v", bifrostErr.Error) + } + if bifrostErr.Error.Code == nil || *bifrostErr.Error.Code != "bad_count" { + t.Fatalf("code mismatch: %#v", bifrostErr.Error) + } + assertNoGigaChatSecretLeak(t, bifrostErr.String()) +} + +func testGigaChatCountTokensRefreshesTokenAfterUnauthorized(t *testing.T) { + t.Parallel() + + var tokenRequests atomic.Int32 + var countTokensRequests atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) { + switch request.URL.Path { + case "/oauth": + tokenIndex := tokenRequests.Add(1) + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(fmt.Sprintf(`{"access_token":"count-tokens-token-%d","expires_at":1893456000}`, tokenIndex))) + case "/v1/tokens/count": + countTokensIndex := countTokensRequests.Add(1) + wantAuthorization := fmt.Sprintf("Bearer count-tokens-token-%d", countTokensIndex) + if got := request.Header.Get("Authorization"); got != wantAuthorization { + t.Fatalf("authorization header mismatch on request %d: got %q, want %q", countTokensIndex, got, wantAuthorization) + } + w.Header().Set("Content-Type", "application/json") + if countTokensIndex == 1 { + w.WriteHeader(http.StatusUnauthorized) + _, _ = w.Write([]byte(`{"status":401,"message":"expired token"}`)) + return + } + _, _ = w.Write([]byte(`[{"tokens":5,"characters":11}]`)) + default: + t.Fatalf("unexpected path: %s", request.URL.Path) + } + })) + defer server.Close() + + provider := newTestGigaChatCountTokensProvider(t, server.URL) + response, bifrostErr := provider.CountTokens(testBifrostContext(), testGigaChatOAuthKey(server.URL+"/oauth", "", "test-credentials"), testGigaChatCountTokensRequest()) + if bifrostErr != nil { + t.Fatalf("CountTokens returned error: %v", bifrostErr) + } + if response == nil || response.InputTokens != 5 { + t.Fatalf("unexpected response: %#v", response) + } + if tokenRequests.Load() != 2 { + t.Fatalf("token request count mismatch: got %d, want 2", tokenRequests.Load()) + } + if countTokensRequests.Load() != 2 { + t.Fatalf("count tokens request count mismatch: got %d, want 2", countTokensRequests.Load()) + } +} + +func testGigaChatCountTokensRequest() *schemas.BifrostResponsesRequest { + return &schemas.BifrostResponsesRequest{ + Model: "GigaChat", + Input: []schemas.ResponsesMessage{ + { + Role: schemas.Ptr(schemas.ResponsesInputMessageRoleUser), + Content: &schemas.ResponsesMessageContent{ContentStr: schemas.Ptr("first")}, + }, + { + Role: schemas.Ptr(schemas.ResponsesInputMessageRoleUser), + Content: &schemas.ResponsesMessageContent{ContentStr: schemas.Ptr("second")}, + }, + }, + } +} + +func newTestGigaChatCountTokensProvider(t *testing.T, baseURL string) *GigaChatProvider { + t.Helper() + + provider := newTestGigaChatChatProvider(t, baseURL) + dialer := &net.Dialer{} + provider.client.Dial = func(addr string) (net.Conn, error) { + return dialer.Dial("tcp", addr) + } + provider.client.DialTimeout = nil + return provider +} + +func assertGigaChatCountTokensHTTPShape(t *testing.T, request *http.Request, wantAuthorization string) { + t.Helper() + + if request.Method != http.MethodPost { + t.Fatalf("method mismatch: got %s", request.Method) + } + if contentType := request.Header.Get("Content-Type"); !strings.Contains(contentType, "application/json") { + t.Fatalf("content type mismatch: got %q", contentType) + } + if accept := request.Header.Get("Accept"); accept != "application/json" { + t.Fatalf("accept mismatch: got %q", accept) + } + if userAgent := request.Header.Get("User-Agent"); userAgent != gigaChatUserAgent { + t.Fatalf("user-agent mismatch: got %q", userAgent) + } + if auth := request.Header.Get("Authorization"); auth != wantAuthorization { + t.Fatalf("authorization mismatch: got %q, want %q", auth, wantAuthorization) + } +} + +func assertGigaChatCountTokensRequestBody(t *testing.T, request *http.Request, wantInput []string) { + t.Helper() + + var body GigaChatCountTokensRequest + if err := json.NewDecoder(request.Body).Decode(&body); err != nil { + t.Fatalf("failed to decode count tokens request body: %v", err) + } + if body.Model != "GigaChat" { + t.Fatalf("model mismatch: got %q", body.Model) + } + if !equalStringSlices(body.Input, wantInput) { + t.Fatalf("input mismatch: got %#v, want %#v", body.Input, wantInput) + } +} + +func equalStringSlices(got []string, want []string) bool { + if len(got) != len(want) { + return false + } + for index := range got { + if got[index] != want[index] { + return false + } + } + return true +} + +func equalIntSlices(got []int, want []int) bool { + if len(got) != len(want) { + return false + } + for index := range got { + if got[index] != want[index] { + return false + } + } + return true +} diff --git a/core/providers/gigachat/embedding_test.go b/core/providers/gigachat/embedding_test.go new file mode 100644 index 00000000000..ced29158f0f --- /dev/null +++ b/core/providers/gigachat/embedding_test.go @@ -0,0 +1,358 @@ +package gigachat + +import ( + "encoding/base64" + "encoding/binary" + "encoding/json" + "fmt" + "math" + "net/http" + "net/http/httptest" + "strings" + "sync/atomic" + "testing" + + "github.com/maximhq/bifrost/core/schemas" +) + +func testGigaChatEmbedding(t *testing.T) { + t.Parallel() + + t.Run("ConverterMapsStringInput", testGigaChatEmbeddingConverterMapsStringInput) + t.Run("ConverterMapsArrayInput", testGigaChatEmbeddingConverterMapsArrayInput) + t.Run("ConverterAcceptsEncodingFormat", testGigaChatEmbeddingConverterAcceptsEncodingFormat) + t.Run("ResponseAppliesBase64EncodingFormat", testGigaChatEmbeddingResponseAppliesBase64EncodingFormat) + t.Run("RejectsUnsupportedParams", testGigaChatEmbeddingRejectsUnsupportedParams) + t.Run("ExecutesWithOAuthToken", testGigaChatEmbeddingExecutesWithOAuthToken) + t.Run("MapsProviderErrors", testGigaChatEmbeddingMapsProviderErrors) + t.Run("RefreshesTokenAfterUnauthorized", testGigaChatEmbeddingRefreshesTokenAfterUnauthorized) +} + +func TestGigaChatEmbedding(t *testing.T) { + testGigaChatEmbedding(t) +} + +func testGigaChatEmbeddingConverterMapsStringInput(t *testing.T) { + t.Parallel() + + text := "hello" + request := &schemas.BifrostEmbeddingRequest{ + Model: "Embeddings", + Input: &schemas.EmbeddingInput{Text: &text}, + } + + gigaChatReq, err := ToGigaChatEmbeddingRequest(request) + if err != nil { + t.Fatalf("ToGigaChatEmbeddingRequest returned error: %v", err) + } + if gigaChatReq.Model != "Embeddings" { + t.Fatalf("model mismatch: got %q", gigaChatReq.Model) + } + + body, err := json.Marshal(gigaChatReq) + if err != nil { + t.Fatalf("failed to marshal request: %v", err) + } + if !strings.Contains(string(body), `"input":"hello"`) { + t.Fatalf("request body should preserve string input, got %s", body) + } +} + +func testGigaChatEmbeddingConverterMapsArrayInput(t *testing.T) { + t.Parallel() + + request := &schemas.BifrostEmbeddingRequest{ + Model: "EmbeddingsGigaR", + Input: &schemas.EmbeddingInput{Texts: []string{"first", "second"}}, + } + + gigaChatReq, err := ToGigaChatEmbeddingRequest(request) + if err != nil { + t.Fatalf("ToGigaChatEmbeddingRequest returned error: %v", err) + } + if gigaChatReq.Model != "EmbeddingsGigaR" { + t.Fatalf("model mismatch: got %q", gigaChatReq.Model) + } + + body, err := json.Marshal(gigaChatReq) + if err != nil { + t.Fatalf("failed to marshal request: %v", err) + } + if !strings.Contains(string(body), `"input":["first","second"]`) { + t.Fatalf("request body should preserve array input, got %s", body) + } +} + +func testGigaChatEmbeddingConverterAcceptsEncodingFormat(t *testing.T) { + t.Parallel() + + encodingFormat := "base64" + request := testGigaChatEmbeddingRequest() + request.Params = &schemas.EmbeddingParameters{ + EncodingFormat: &encodingFormat, + } + + gigaChatReq, err := ToGigaChatEmbeddingRequest(request) + if err != nil { + t.Fatalf("ToGigaChatEmbeddingRequest returned error: %v", err) + } + + body, err := json.Marshal(gigaChatReq) + if err != nil { + t.Fatalf("failed to marshal request: %v", err) + } + if strings.Contains(string(body), "encoding_format") { + t.Fatalf("GigaChat request body should not include encoding_format, got %s", body) + } +} + +func testGigaChatEmbeddingResponseAppliesBase64EncodingFormat(t *testing.T) { + t.Parallel() + + encodingFormat := "base64" + response := ToBifrostEmbeddingResponse(schemas.GigaChat, &GigaChatEmbeddingResponse{ + Object: "list", + Model: "Embeddings", + Data: []GigaChatEmbeddingData{{ + Object: "embedding", + Index: 0, + Embedding: []float64{0.1, 0.2}, + }}, + }) + + if err := applyGigaChatEmbeddingEncodingFormat(response, &schemas.EmbeddingParameters{EncodingFormat: &encodingFormat}); err != nil { + t.Fatalf("applyGigaChatEmbeddingEncodingFormat returned error: %v", err) + } + if response.Data[0].Embedding.EmbeddingArray != nil { + t.Fatalf("expected base64 embedding string, got float array %#v", response.Data[0].Embedding.EmbeddingArray) + } + if response.Data[0].Embedding.EmbeddingStr == nil { + t.Fatal("expected base64 embedding string, got nil") + } + + decoded, err := base64.StdEncoding.DecodeString(*response.Data[0].Embedding.EmbeddingStr) + if err != nil { + t.Fatalf("failed to decode base64 embedding: %v", err) + } + if len(decoded) != 8 { + t.Fatalf("decoded embedding byte length mismatch: got %d, want 8", len(decoded)) + } + gotFirst := math.Float32frombits(binary.LittleEndian.Uint32(decoded[0:4])) + gotSecond := math.Float32frombits(binary.LittleEndian.Uint32(decoded[4:8])) + if gotFirst != float32(0.1) || gotSecond != float32(0.2) { + t.Fatalf("decoded embedding mismatch: got [%v %v]", gotFirst, gotSecond) + } +} + +func testGigaChatEmbeddingRejectsUnsupportedParams(t *testing.T) { + t.Parallel() + + encodingFormat := "base64" + dimensions := 1024 + request := testGigaChatEmbeddingRequest() + request.Params = &schemas.EmbeddingParameters{ + EncodingFormat: &encodingFormat, + Dimensions: &dimensions, + ExtraParams: map[string]interface{}{ + "user": "user-id", + }, + } + + _, err := ToGigaChatEmbeddingRequest(request) + if err == nil { + t.Fatal("expected unsupported params error, got nil") + } + for _, want := range []string{"dimensions", "user"} { + if !strings.Contains(err.Error(), want) { + t.Fatalf("error %q missing unsupported param %q", err.Error(), want) + } + } + if strings.Contains(err.Error(), "encoding_format") { + t.Fatalf("encoding_format should be accepted for OpenAI SDK compatibility, got %q", err.Error()) + } + + _, err = ToGigaChatEmbeddingRequest(&schemas.BifrostEmbeddingRequest{ + Model: "Embeddings", + Input: &schemas.EmbeddingInput{Embedding: []int{1, 2, 3}}, + }) + if err == nil || !strings.Contains(err.Error(), "string or array-of-string") { + t.Fatalf("expected unsupported input error, got %v", err) + } +} + +func testGigaChatEmbeddingExecutesWithOAuthToken(t *testing.T) { + t.Parallel() + + var tokenRequests atomic.Int32 + var embeddingRequests atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) { + switch request.URL.Path { + case "/oauth": + tokenRequests.Add(1) + if got := request.Header.Get("Authorization"); got != "Basic super-secret-credentials" { + t.Fatalf("token authorization header mismatch: got %q", got) + } + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"access_token":"embedding-access-token","expires_at":1893456000}`)) + case "/v1/embeddings": + embeddingRequests.Add(1) + if request.Method != http.MethodPost { + t.Fatalf("method mismatch: got %s, want POST", request.Method) + } + if got := request.Header.Get("Authorization"); got != "Bearer embedding-access-token" { + t.Fatalf("embeddings authorization header mismatch: got %q", got) + } + if strings.Contains(request.Header.Get("Authorization"), "super-secret-credentials") { + t.Fatal("embeddings request leaked OAuth credentials") + } + assertGigaChatEmbeddingRequestBody(t, request) + w.Header().Set("Content-Type", "application/json") + w.Header().Set("X-Request-ID", "embeddings-request-id") + _, _ = w.Write([]byte(`{ + "object":"list", + "data":[ + {"object":"embedding","embedding":[0.1,0.2],"index":0,"usage":{"prompt_tokens":5}}, + {"object":"embedding","embedding":[0.3,0.4],"index":1,"usage":{"prompt_tokens":7}} + ], + "model":"Embeddings" + }`)) + default: + t.Fatalf("unexpected path: %s", request.URL.Path) + } + })) + defer server.Close() + + provider := newTestGigaChatChatProvider(t, server.URL) + provider.sendBackRawRequest = true + provider.sendBackRawResponse = true + + response, bifrostErr := provider.Embedding(testBifrostContext(), testGigaChatOAuthKey(server.URL+"/oauth", "", "super-secret-credentials"), testGigaChatEmbeddingRequest()) + if bifrostErr != nil { + t.Fatalf("Embedding returned error: %v", bifrostErr) + } + if tokenRequests.Load() != 1 { + t.Fatalf("token request count mismatch: got %d, want 1", tokenRequests.Load()) + } + if embeddingRequests.Load() != 1 { + t.Fatalf("embedding request count mismatch: got %d, want 1", embeddingRequests.Load()) + } + if response.Model != "Embeddings" || response.Object != "list" { + t.Fatalf("response metadata mismatch: %#v", response) + } + if response.ExtraFields.Provider != schemas.GigaChat { + t.Fatalf("provider mismatch: got %q, want %q", response.ExtraFields.Provider, schemas.GigaChat) + } + if len(response.Data) != 2 { + t.Fatalf("embedding count mismatch: got %d, want 2", len(response.Data)) + } + if got := response.Data[0].Embedding.EmbeddingArray; fmt.Sprint(got) != fmt.Sprint([]float64{0.1, 0.2}) { + t.Fatalf("embedding mismatch: %#v", got) + } + if response.Usage == nil || response.Usage.PromptTokens != 12 || response.Usage.TotalTokens != 12 { + t.Fatalf("usage mismatch: %#v", response.Usage) + } + if response.ExtraFields.RawRequest == nil || response.ExtraFields.RawResponse == nil { + t.Fatalf("expected raw request and response, got request=%#v response=%#v", response.ExtraFields.RawRequest, response.ExtraFields.RawResponse) + } +} + +func testGigaChatEmbeddingMapsProviderErrors(t *testing.T) { + t.Parallel() + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusBadRequest) + _, _ = w.Write([]byte(`{"status":400,"code":123,"message":"bad embeddings request"}`)) + })) + defer server.Close() + + provider := newTestGigaChatChatProvider(t, server.URL) + response, bifrostErr := provider.Embedding(testBifrostContext(), testGigaChatAccessTokenKey("provider-error-token"), testGigaChatEmbeddingRequest()) + if response != nil { + t.Fatalf("expected nil response, got %#v", response) + } + if bifrostErr == nil { + t.Fatal("expected provider error, got nil") + } + if bifrostErr.StatusCode == nil || *bifrostErr.StatusCode != http.StatusBadRequest { + t.Fatalf("status mismatch: %#v", bifrostErr.StatusCode) + } + if bifrostErr.Error == nil || bifrostErr.Error.Message != "bad embeddings request" { + t.Fatalf("message mismatch: %#v", bifrostErr.Error) + } + if bifrostErr.Error.Code == nil || *bifrostErr.Error.Code != "123" { + t.Fatalf("code mismatch: %#v", bifrostErr.Error) + } + assertNoGigaChatSecretLeak(t, bifrostErr.String()) +} + +func testGigaChatEmbeddingRefreshesTokenAfterUnauthorized(t *testing.T) { + t.Parallel() + + var tokenRequests atomic.Int32 + var embeddingRequests atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) { + switch request.URL.Path { + case "/oauth": + tokenIndex := tokenRequests.Add(1) + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(fmt.Sprintf(`{"access_token":"embedding-token-%d","expires_at":1893456000}`, tokenIndex))) + case "/v1/embeddings": + embeddingIndex := embeddingRequests.Add(1) + wantAuthorization := fmt.Sprintf("Bearer embedding-token-%d", embeddingIndex) + if got := request.Header.Get("Authorization"); got != wantAuthorization { + t.Fatalf("authorization header mismatch on request %d: got %q, want %q", embeddingIndex, got, wantAuthorization) + } + w.Header().Set("Content-Type", "application/json") + if embeddingIndex == 1 { + w.WriteHeader(http.StatusUnauthorized) + _, _ = w.Write([]byte(`{"status":401,"message":"expired token"}`)) + return + } + _, _ = w.Write([]byte(`{"object":"list","data":[{"object":"embedding","embedding":[0.1],"index":0}],"model":"Embeddings"}`)) + default: + t.Fatalf("unexpected path: %s", request.URL.Path) + } + })) + defer server.Close() + + provider := newTestGigaChatChatProvider(t, server.URL) + response, bifrostErr := provider.Embedding(testBifrostContext(), testGigaChatOAuthKey(server.URL+"/oauth", "", "test-credentials"), testGigaChatEmbeddingRequest()) + if bifrostErr != nil { + t.Fatalf("Embedding returned error: %v", bifrostErr) + } + if response == nil || len(response.Data) != 1 { + t.Fatalf("unexpected response: %#v", response) + } + if tokenRequests.Load() != 2 { + t.Fatalf("token request count mismatch: got %d, want 2", tokenRequests.Load()) + } + if embeddingRequests.Load() != 2 { + t.Fatalf("embedding request count mismatch: got %d, want 2", embeddingRequests.Load()) + } +} + +func testGigaChatEmbeddingRequest() *schemas.BifrostEmbeddingRequest { + return &schemas.BifrostEmbeddingRequest{ + Model: "Embeddings", + Input: &schemas.EmbeddingInput{Texts: []string{"first", "second"}}, + } +} + +func assertGigaChatEmbeddingRequestBody(t *testing.T, request *http.Request) { + t.Helper() + + var body GigaChatEmbeddingRequest + if err := json.NewDecoder(request.Body).Decode(&body); err != nil { + t.Fatalf("failed to decode embeddings request body: %v", err) + } + if body.Model != "Embeddings" { + t.Fatalf("model mismatch: got %q", body.Model) + } + if body.Input == nil || len(body.Input.Texts) != 2 || body.Input.Texts[0] != "first" || body.Input.Texts[1] != "second" { + t.Fatalf("input mismatch: %#v", body.Input) + } + if body.Input.Text != nil || body.Input.Embedding != nil || body.Input.Embeddings != nil { + t.Fatalf("unexpected non-text embedding input: %#v", body.Input) + } +} diff --git a/core/providers/gigachat/errors_test.go b/core/providers/gigachat/errors_test.go new file mode 100644 index 00000000000..30d0b092494 --- /dev/null +++ b/core/providers/gigachat/errors_test.go @@ -0,0 +1,254 @@ +package gigachat + +import ( + "encoding/json" + "net/http" + "strings" + "testing" + + "github.com/maximhq/bifrost/core/schemas" + "github.com/valyala/fasthttp" +) + +func testGigaChatErrors(t *testing.T) { + t.Parallel() + + t.Run("ParsesCommonPayloads", testGigaChatErrorParsesCommonPayloads) + t.Run("ParsesOAuthPayloads", testGigaChatErrorParsesOAuthPayloads) + t.Run("UsesFallbackForNonJSON", testGigaChatErrorUsesFallbackForNonJSON) + t.Run("RedactsRawPayloads", testGigaChatErrorRedactsRawPayloads) + t.Run("RedactsExpandedRawAuthMaterial", testGigaChatErrorRedactsExpandedRawAuthMaterial) + t.Run("RedactsExistingRawResponse", testGigaChatErrorRedactsExistingRawResponse) + t.Run("RedactsStreamingCallbackRawResponse", testGigaChatErrorRedactsStreamingCallbackRawResponse) + t.Run("PreservesSafeRawPayloadOrder", testGigaChatErrorPreservesSafeRawPayloadOrder) + t.Run("RedactsTextPayloads", testGigaChatErrorRedactsTextPayloads) +} + +func TestGigaChatErrors(t *testing.T) { + testGigaChatErrors(t) +} + +func testGigaChatErrorParsesCommonPayloads(t *testing.T) { + t.Parallel() + + resp := fasthttp.AcquireResponse() + defer fasthttp.ReleaseResponse(resp) + resp.SetStatusCode(http.StatusTooManyRequests) + resp.SetBodyString(`{"status":429,"code":7,"message":"quota exceeded"}`) + + bifrostErr := ParseGigaChatError(resp, schemas.GigaChat) + if bifrostErr == nil { + t.Fatal("expected error, got nil") + } + if bifrostErr.StatusCode == nil || *bifrostErr.StatusCode != http.StatusTooManyRequests { + t.Fatalf("status mismatch: %#v", bifrostErr.StatusCode) + } + if bifrostErr.Error == nil || bifrostErr.Error.Message != "quota exceeded" { + t.Fatalf("message mismatch: %#v", bifrostErr.Error) + } + if bifrostErr.Error.Code == nil || *bifrostErr.Error.Code != "7" { + t.Fatalf("code mismatch: %#v", bifrostErr.Error) + } + if bifrostErr.ExtraFields.Provider != schemas.GigaChat { + t.Fatalf("provider mismatch: got %q", bifrostErr.ExtraFields.Provider) + } +} + +func testGigaChatErrorParsesOAuthPayloads(t *testing.T) { + t.Parallel() + + resp := fasthttp.AcquireResponse() + defer fasthttp.ReleaseResponse(resp) + resp.SetStatusCode(http.StatusUnauthorized) + resp.SetBodyString(`{"error":"invalid_client","error_description":"bad credentials","code":"AUTH_FAILED"}`) + + bifrostErr := ParseGigaChatError(resp, schemas.GigaChat) + if bifrostErr == nil { + t.Fatal("expected error, got nil") + } + if bifrostErr.Error == nil || bifrostErr.Error.Message != "bad credentials" { + t.Fatalf("message mismatch: %#v", bifrostErr.Error) + } + if bifrostErr.Error.Code == nil || *bifrostErr.Error.Code != "AUTH_FAILED" { + t.Fatalf("code mismatch: %#v", bifrostErr.Error) + } +} + +func testGigaChatErrorUsesFallbackForNonJSON(t *testing.T) { + t.Parallel() + + resp := fasthttp.AcquireResponse() + defer fasthttp.ReleaseResponse(resp) + resp.SetStatusCode(http.StatusBadGateway) + resp.SetBodyString("upstream unavailable") + + bifrostErr := ParseGigaChatError(resp, schemas.GigaChat) + if bifrostErr == nil { + t.Fatal("expected error, got nil") + } + if bifrostErr.Error == nil || !strings.Contains(bifrostErr.Error.Message, "provider API error") { + t.Fatalf("message mismatch: %#v", bifrostErr.Error) + } +} + +func testGigaChatErrorRedactsRawPayloads(t *testing.T) { + t.Parallel() + + ctx := testBifrostContext() + requestBody := []byte(`{"model":"GigaChat","credentials":"super-secret-credentials","nested":{"password":"super-secret-password"}}`) + responseBody := []byte(`{"message":"bad","access_token":"super-secret-token","authorization":"Bearer super-secret-token"}`) + bifrostErr := newGigaChatProviderResponseError("failed", nil) + + enriched := enrichGigaChatError(ctx, bifrostErr, requestBody, responseBody, true, true) + output := stringifyGigaChatRaw(enriched.ExtraFields.RawRequest) + stringifyGigaChatRaw(enriched.ExtraFields.RawResponse) + for _, secret := range []string{"super-secret-credentials", "super-secret-password", "super-secret-token"} { + if strings.Contains(output, secret) { + t.Fatalf("raw payload leaked %q in %s", secret, output) + } + } + if !strings.Contains(output, "redacted") { + t.Fatalf("expected redacted marker in raw payloads, got %s", output) + } +} + +func testGigaChatErrorRedactsTextPayloads(t *testing.T) { + t.Parallel() + + payload := []byte("error: authorization bearer super-secret-token failed; access_token=super-secret-access; user=super-secret-user; password=super-secret-password; key_file=/secure/client.key; Basic super-secret-basic rejected; -----BEGIN PRIVATE KEY-----\nsuper-secret-private-key\n-----END PRIVATE KEY-----") + redacted := string(redactGigaChatRawPayload(payload)) + for _, secret := range []string{"super-secret-token", "super-secret-access", "super-secret-user", "super-secret-password", "/secure/client.key", "super-secret-basic", "super-secret-private-key"} { + if strings.Contains(redacted, secret) { + t.Fatalf("text payload leaked %q in %s", secret, redacted) + } + } +} + +func testGigaChatErrorRedactsExpandedRawAuthMaterial(t *testing.T) { + t.Parallel() + + ctx := testBifrostContext() + requestBody := []byte(`{ + "model":"GigaChat", + "authorization":"Basic super-secret-request-basic", + "credentials":"super-secret-credentials", + "user":"super-secret-user", + "password":"super-secret-password", + "key_file":"/secure/client.key", + "cert_file":"/secure/client.crt", + "ca_bundle_file":"/secure/ca.crt", + "private_key":"-----BEGIN PRIVATE KEY-----\nsuper-secret-private-key\n-----END PRIVATE KEY-----", + "messages":[{"role":"user","content":"safe prompt"}] + }`) + responseBody := []byte(`{ + "message":"bad", + "authorization":"Bearer super-secret-response-bearer", + "access_token":"super-secret-access-token", + "client_secret":"super-secret-client-secret", + "refresh_token":"super-secret-refresh-token", + "errors":["Basic super-secret-array-basic"] + }`) + bifrostErr := newGigaChatProviderResponseError("authorization Bearer super-secret-error-bearer failed with password=super-secret-error-password and -----BEGIN PRIVATE KEY-----\nsuper-secret-error-private-key\n-----END PRIVATE KEY-----", nil) + + enriched := enrichGigaChatError(ctx, bifrostErr, requestBody, responseBody, true, true) + requestOutput := stringifyGigaChatRaw(enriched.ExtraFields.RawRequest) + responseOutput := stringifyGigaChatRaw(enriched.ExtraFields.RawResponse) + errorOutput := enriched.String() + + assertGigaChatOutputOmits(t, "raw request", requestOutput, []string{ + "super-secret-request-basic", + "super-secret-credentials", + "super-secret-user", + "super-secret-password", + "/secure/client.key", + "/secure/client.crt", + "/secure/ca.crt", + "super-secret-private-key", + }) + assertGigaChatOutputOmits(t, "raw response", responseOutput, []string{ + "super-secret-response-bearer", + "super-secret-access-token", + "super-secret-client-secret", + "super-secret-refresh-token", + "super-secret-array-basic", + }) + assertGigaChatOutputOmits(t, "error message", errorOutput, []string{ + "super-secret-error-bearer", + "super-secret-error-password", + "super-secret-error-private-key", + }) + if !strings.Contains(requestOutput+responseOutput+errorOutput, "redacted") { + t.Fatalf("expected redacted markers, got request=%s response=%s error=%s", requestOutput, responseOutput, errorOutput) + } +} + +func testGigaChatErrorRedactsExistingRawResponse(t *testing.T) { + t.Parallel() + + ctx := testBifrostContext() + resp := fasthttp.AcquireResponse() + defer fasthttp.ReleaseResponse(resp) + resp.SetStatusCode(http.StatusBadRequest) + resp.SetBodyString(`{"status":400,"message":"authorization bearer super-secret-token","access_token":"super-secret-token"}`) + + bifrostErr := ParseGigaChatError(resp, schemas.GigaChat) + enriched := enrichGigaChatError(ctx, bifrostErr, nil, nil, false, true) + output := enriched.String() + stringifyGigaChatRaw(enriched.ExtraFields.RawResponse) + if strings.Contains(output, "super-secret-token") { + t.Fatalf("existing raw response leaked secret in %s", output) + } + if !strings.Contains(output, "redacted") { + t.Fatalf("expected redacted marker in existing raw response, got %s", output) + } +} + +func testGigaChatErrorRedactsStreamingCallbackRawResponse(t *testing.T) { + t.Parallel() + + handler := handleGigaChatChatStreamResponse(schemas.GigaChat) + var response schemas.BifrostChatResponse + _, rawResponse, bifrostErr := handler( + []byte(`{"status":401,"message":"bearer super-secret-token","access_token":"super-secret-token"}`), + &response, + []byte(`{"model":"GigaChat"}`), + true, + true, + ) + if bifrostErr == nil { + t.Fatal("expected streaming error, got nil") + } + + output := bifrostErr.String() + stringifyGigaChatRaw(rawResponse) + if strings.Contains(output, "super-secret-token") { + t.Fatalf("streaming raw response leaked secret in %s", output) + } +} + +func testGigaChatErrorPreservesSafeRawPayloadOrder(t *testing.T) { + t.Parallel() + + payload := []byte(`{"z":1,"a":2,"message":"safe"}`) + redacted := redactGigaChatRawPayload(payload) + if string(redacted) != string(payload) { + t.Fatalf("safe payload order changed: got %s, want %s", redacted, payload) + } +} + +func stringifyGigaChatRaw(raw interface{}) string { + if raw == nil { + return "" + } + data, err := json.Marshal(raw) + if err != nil { + return "" + } + return string(data) +} + +func assertGigaChatOutputOmits(t *testing.T, label string, output string, secrets []string) { + t.Helper() + for _, secret := range secrets { + if strings.Contains(output, secret) { + t.Fatalf("%s leaked %q in %s", label, secret, output) + } + } +} diff --git a/core/providers/gigachat/files_test.go b/core/providers/gigachat/files_test.go new file mode 100644 index 00000000000..8ccf2f35ae9 --- /dev/null +++ b/core/providers/gigachat/files_test.go @@ -0,0 +1,739 @@ +package gigachat + +import ( + "bytes" + "encoding/base64" + "encoding/json" + "io" + "net/http" + "net/http/httptest" + "strings" + "sync/atomic" + "testing" + + schemas "github.com/maximhq/bifrost/core/schemas" +) + +func TestToGigaChatFilePurpose(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + purpose schemas.FilePurpose + want string + }{ + {name: "assistants maps to assistant", purpose: schemas.FilePurposeAssistants, want: gigaChatFilePurposeAssistant}, + {name: "batch maps to general", purpose: schemas.FilePurposeBatch, want: gigaChatFilePurposeGeneral}, + {name: "fine tune maps to general", purpose: schemas.FilePurposeFineTune, want: gigaChatFilePurposeGeneral}, + {name: "vision maps to general", purpose: schemas.FilePurposeVision, want: gigaChatFilePurposeGeneral}, + {name: "responses maps to general", purpose: schemas.FilePurposeResponses, want: gigaChatFilePurposeGeneral}, + {name: "evals maps to general", purpose: schemas.FilePurposeEvals, want: gigaChatFilePurposeGeneral}, + {name: "user data maps to general", purpose: schemas.FilePurposeUserData, want: gigaChatFilePurposeGeneral}, + {name: "batch output maps to general", purpose: schemas.FilePurposeBatchOutput, want: gigaChatFilePurposeGeneral}, + {name: "empty maps to general", purpose: "", want: gigaChatFilePurposeGeneral}, + } + + for _, tt := range tests { + tt := tt + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + if got := toGigaChatFilePurpose(tt.purpose); got != tt.want { + t.Fatalf("toGigaChatFilePurpose(%q) = %q, want %q", tt.purpose, got, tt.want) + } + }) + } +} + +func TestToBifrostFilePurpose(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + gigaChatPurpose string + requestedPurpose schemas.FilePurpose + want schemas.FilePurpose + }{ + { + name: "assistant maps to assistants", + gigaChatPurpose: gigaChatFilePurposeAssistant, + want: schemas.FilePurposeAssistants, + }, + { + name: "general defaults to user data", + gigaChatPurpose: gigaChatFilePurposeGeneral, + want: schemas.FilePurposeUserData, + }, + { + name: "general keeps requested purpose when available", + gigaChatPurpose: gigaChatFilePurposeGeneral, + requestedPurpose: schemas.FilePurposeBatch, + want: schemas.FilePurposeBatch, + }, + { + name: "empty defaults to user data", + want: schemas.FilePurposeUserData, + }, + { + name: "unknown purpose is preserved", + gigaChatPurpose: "custom_purpose", + want: schemas.FilePurpose("custom_purpose"), + }, + { + name: "general trims whitespace", + gigaChatPurpose: " general ", + want: schemas.FilePurposeUserData, + }, + } + + for _, tt := range tests { + tt := tt + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + got := toBifrostFilePurpose(tt.gigaChatPurpose, tt.requestedPurpose) + if got != tt.want { + t.Fatalf("toBifrostFilePurpose(%q, %q) = %q, want %q", tt.gigaChatPurpose, tt.requestedPurpose, got, tt.want) + } + }) + } +} + +func TestGigaChatFileTypesJSON(t *testing.T) { + t.Parallel() + + accessPolicy := "public" + uploaded := GigaChatUploadedFile{ + ID: "file-1", + Object: "file", + Bytes: 123, + CreatedAt: 1780306293, + Filename: "document.txt", + Purpose: gigaChatFilePurposeGeneral, + AccessPolicy: &accessPolicy, + } + + raw, err := json.Marshal(GigaChatUploadedFiles{Data: []GigaChatUploadedFile{uploaded}}) + if err != nil { + t.Fatalf("marshal uploaded files: %v", err) + } + + var decoded GigaChatUploadedFiles + if err := json.Unmarshal(raw, &decoded); err != nil { + t.Fatalf("unmarshal uploaded files: %v", err) + } + if len(decoded.Data) != 1 { + t.Fatalf("decoded %d files, want 1", len(decoded.Data)) + } + if decoded.Data[0].ID != uploaded.ID || decoded.Data[0].Purpose != uploaded.Purpose { + t.Fatalf("decoded file = %+v, want %+v", decoded.Data[0], uploaded) + } + + contentRaw, err := json.Marshal(GigaChatFileContent{Content: "SGVsbG8="}) + if err != nil { + t.Fatalf("marshal file content: %v", err) + } + if string(contentRaw) != `{"content":"SGVsbG8="}` { + t.Fatalf("content JSON = %s", contentRaw) + } +} + +func TestGigaChatFilesHTTP(t *testing.T) { + t.Parallel() + + t.Run("UploadMultipart", testGigaChatFileUploadMultipart) + t.Run("ListUsesKeyBaseURLAndAuthHeaders", testGigaChatFileListUsesKeyBaseURLAndAuthHeaders) + t.Run("ListReturnsAllFilesAndRawRequest", testGigaChatFileListReturnsAllFilesAndRawRequest) + t.Run("ListPaginatesLocally", testGigaChatFileListPaginatesLocally) + t.Run("ListRejectsUnsupportedOrder", testGigaChatFileListRejectsUnsupportedOrder) + t.Run("ListPreservesUpstreamPurposeWhenFiltering", testGigaChatFileListPreservesUpstreamPurposeWhenFiltering) + t.Run("ListRetrieveDelete", testGigaChatFileListRetrieveDelete) + t.Run("ContentRawBytes", testGigaChatFileContentRawBytes) + t.Run("ContentBase64Wrapper", testGigaChatFileContentBase64Wrapper) + t.Run("RefreshesTokenAfterUnauthorized", testGigaChatFileUploadRefreshesTokenAfterUnauthorized) +} + +func testGigaChatFileUploadMultipart(t *testing.T) { + t.Parallel() + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) { + if request.URL.Path != "/v1/files" { + t.Fatalf("path mismatch: got %s", request.URL.Path) + } + if request.Method != http.MethodPost { + t.Fatalf("method mismatch: got %s, want POST", request.Method) + } + if got := request.Header.Get("Authorization"); got != "Bearer file-upload-token" { + t.Fatalf("authorization header mismatch: got %q", got) + } + if !strings.HasPrefix(request.Header.Get("Content-Type"), "multipart/form-data;") { + t.Fatalf("content type mismatch: got %q", request.Header.Get("Content-Type")) + } + if err := request.ParseMultipartForm(1024); err != nil { + t.Fatalf("ParseMultipartForm returned error: %v", err) + } + if got := request.FormValue("purpose"); got != gigaChatFilePurposeGeneral { + t.Fatalf("purpose mismatch: got %q, want %q", got, gigaChatFilePurposeGeneral) + } + file, header, err := request.FormFile("file") + if err != nil { + t.Fatalf("FormFile returned error: %v", err) + } + defer file.Close() + if header.Filename != "input.txt" { + t.Fatalf("filename mismatch: got %q", header.Filename) + } + if got := header.Header.Get("Content-Type"); got != "text/plain" { + t.Fatalf("file content type mismatch: got %q, want text/plain", got) + } + body, err := io.ReadAll(file) + if err != nil { + t.Fatalf("ReadAll returned error: %v", err) + } + if !bytes.Equal(body, []byte(`{"ok":true}`)) { + t.Fatalf("file body mismatch: got %q", string(body)) + } + + w.Header().Set("Content-Type", "application/json") + w.Header().Set("X-Request-ID", "file-upload-request-id") + _, _ = w.Write([]byte(`{"id":"file-uploaded","object":"file","bytes":11,"created_at":1780306293,"filename":"input.txt","purpose":"general","access_policy":"private"}`)) + })) + defer server.Close() + + provider := newTestGigaChatChatProvider(t, server.URL) + key := testGigaChatAccessTokenKey("file-upload-token") + contentType := "application/jsonl" + + response, bifrostErr := provider.FileUpload(testBifrostContext(), key, &schemas.BifrostFileUploadRequest{ + Provider: schemas.GigaChat, + File: []byte(`{"ok":true}`), + Filename: "input.jsonl", + Purpose: schemas.FilePurposeBatch, + ContentType: &contentType, + }) + if bifrostErr != nil { + t.Fatalf("FileUpload returned error: %v", bifrostErr) + } + if response.ID != "file-uploaded" || response.Filename != "input.txt" || response.Bytes != 11 { + t.Fatalf("unexpected upload response: %#v", response) + } + if response.Purpose != schemas.FilePurposeBatch { + t.Fatalf("purpose mismatch: got %q, want %q", response.Purpose, schemas.FilePurposeBatch) + } + if response.StorageBackend != schemas.FileStorageAPI { + t.Fatalf("storage backend mismatch: got %q", response.StorageBackend) + } + if response.ExtraFields.ProviderResponseHeaders["X-Request-Id"] != "file-upload-request-id" { + t.Fatalf("provider headers mismatch: %#v", response.ExtraFields.ProviderResponseHeaders) + } +} + +func TestGigaChatFileUploadMetadata(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + filename string + contentType *string + file []byte + wantFilename string + wantContentType string + }{ + { + name: "json content type is uploaded as supported text file", + filename: "input.jsonl", + contentType: schemas.Ptr("application/json"), + file: []byte(`{"custom_id":"req-1"}` + "\n"), + wantFilename: "input.txt", + wantContentType: "text/plain", + }, + { + name: "missing filename text defaults to txt", + file: []byte(`{"custom_id":"req-1"}` + "\n"), + wantFilename: "file.txt", + wantContentType: "text/plain", + }, + { + name: "supported binary type is preserved", + filename: "document.pdf", + contentType: schemas.Ptr("application/pdf"), + file: []byte("%PDF-1.7\n"), + wantFilename: "document.pdf", + wantContentType: "application/pdf", + }, + { + name: "xlsx mime alias uses gigachat supported type", + filename: "table.xlsx", + contentType: schemas.Ptr("application/vnd.openxmlformats-officedocument.spreadsheetml.sheet"), + file: []byte("xlsx"), + wantFilename: "table.xlsx", + wantContentType: "application/vnd.ms-excel", + }, + } + + for _, tt := range tests { + tt := tt + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + body, contentType, bifrostErr := buildGigaChatFileUploadBody(&schemas.BifrostFileUploadRequest{ + Provider: schemas.GigaChat, + File: tt.file, + Filename: tt.filename, + Purpose: schemas.FilePurposeUserData, + ContentType: tt.contentType, + }) + if bifrostErr != nil { + t.Fatalf("buildGigaChatFileUploadBody returned error: %v", bifrostErr) + } + if !strings.HasPrefix(contentType, "multipart/form-data;") { + t.Fatalf("content type mismatch: got %q", contentType) + } + + request := httptest.NewRequest(http.MethodPost, "/files", bytes.NewReader(body)) + request.Header.Set("Content-Type", contentType) + if err := request.ParseMultipartForm(1024); err != nil { + t.Fatalf("ParseMultipartForm returned error: %v", err) + } + file, header, err := request.FormFile("file") + if err != nil { + t.Fatalf("FormFile returned error: %v", err) + } + defer file.Close() + if header.Filename != tt.wantFilename { + t.Fatalf("filename mismatch: got %q, want %q", header.Filename, tt.wantFilename) + } + if got := header.Header.Get("Content-Type"); got != tt.wantContentType { + t.Fatalf("file content type mismatch: got %q, want %q", got, tt.wantContentType) + } + }) + } +} + +func TestGigaChatMultipartFilenameEscaping(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + filename string + want string + }{ + { + name: "normal unicode filename is preserved", + filename: "отчет.txt", + want: "отчет.txt", + }, + { + name: "quotes are escaped", + filename: `report "final".txt`, + want: `report \"final\".txt`, + }, + { + name: "backslashes are escaped", + filename: `dir\file.txt`, + want: `dir\\file.txt`, + }, + { + name: "line feed is replaced", + filename: "bad\nname.txt", + want: "bad_name.txt", + }, + { + name: "carriage return is replaced", + filename: "bad\rname.txt", + want: "bad_name.txt", + }, + { + name: "crlf is replaced", + filename: "bad\r\nname.txt", + want: "bad__name.txt", + }, + { + name: "other control characters are replaced", + filename: "bad\x00\t\x7fname.txt", + want: "bad___name.txt", + }, + { + name: "sanitized filename can still escape quotes and backslashes", + filename: "bad\r\ndir\\file \"final\".txt", + want: "bad__dir\\\\file \\\"final\\\".txt", + }, + } + + for _, tt := range tests { + tt := tt + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + got := escapeGigaChatMultipartFilename(tt.filename) + if got != tt.want { + t.Fatalf("escapeGigaChatMultipartFilename(%q) = %q, want %q", tt.filename, got, tt.want) + } + if strings.ContainsAny(got, "\r\n") { + t.Fatalf("escaped filename still contains CR/LF: %q", got) + } + }) + } +} + +func testGigaChatFileListUsesKeyBaseURLAndAuthHeaders(t *testing.T) { + t.Parallel() + + networkServer := httptest.NewServer(http.HandlerFunc(func(_ http.ResponseWriter, request *http.Request) { + t.Fatalf("network base_url server should not be used, got %s", request.URL.Path) + })) + defer networkServer.Close() + + keyServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) { + if request.URL.Path != "/custom-api/v1/files" { + t.Fatalf("path mismatch: got %s", request.URL.Path) + } + if request.Method != http.MethodGet { + t.Fatalf("method mismatch: got %s, want GET", request.Method) + } + if got := request.Header.Get("Authorization"); got != "Bearer key-base-url-token" { + t.Fatalf("authorization header mismatch: got %q", got) + } + if got := request.Header.Get(gigaChatUserAgentHeader); got != gigaChatUserAgent { + t.Fatalf("user-agent mismatch: got %q", got) + } + + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"data":[]}`)) + })) + defer keyServer.Close() + + provider := newTestGigaChatChatProvider(t, networkServer.URL) + key := schemas.Key{ + GigaChatKeyConfig: &schemas.GigaChatKeyConfig{ + AccessToken: schemas.NewSecretVar("key-base-url-token"), + BaseURL: keyServer.URL + "/custom-api", + }, + } + + response, bifrostErr := provider.FileList(testBifrostContext(), []schemas.Key{key}, &schemas.BifrostFileListRequest{Provider: schemas.GigaChat}) + if bifrostErr != nil { + t.Fatalf("FileList returned error: %v", bifrostErr) + } + if response.Object != "list" || len(response.Data) != 0 { + t.Fatalf("unexpected list response: %#v", response) + } +} + +func testGigaChatFileListReturnsAllFilesAndRawRequest(t *testing.T) { + t.Parallel() + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) { + if request.URL.Path != "/v1/files" { + t.Fatalf("path mismatch: got %s", request.URL.Path) + } + if request.Method != http.MethodGet { + t.Fatalf("method mismatch: got %s, want GET", request.Method) + } + + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"data":[{"id":"file-1","object":"file","bytes":10,"created_at":1780306293,"filename":"one.txt","purpose":"general"},{"id":"file-2","object":"file","bytes":20,"created_at":1780306294,"filename":"two.txt","purpose":"general"}]}`)) + })) + defer server.Close() + + provider := newTestGigaChatChatProvider(t, server.URL) + provider.sendBackRawRequest = true + provider.sendBackRawResponse = true + + ctx := testBifrostContext() + ctx.SetValue(schemas.BifrostContextKeyCaptureRawRequest, true) + ctx.SetValue(schemas.BifrostContextKeyCaptureRawResponse, true) + + response, bifrostErr := provider.FileList(ctx, []schemas.Key{testGigaChatAccessTokenKey("files-token")}, &schemas.BifrostFileListRequest{ + Provider: schemas.GigaChat, + Limit: 10, + }) + if bifrostErr != nil { + t.Fatalf("FileList returned error: %v", bifrostErr) + } + if len(response.Data) != 2 || response.Data[0].ID != "file-1" || response.Data[1].ID != "file-2" { + t.Fatalf("file list mismatch: %#v", response.Data) + } + if got := stringifyGigaChatRaw(response.ExtraFields.RawRequest); got != `{}` { + t.Fatalf("raw request mismatch: got %s", got) + } + if response.ExtraFields.RawResponse == nil { + t.Fatal("expected raw response") + } +} + +func testGigaChatFileListPaginatesLocally(t *testing.T) { + t.Parallel() + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) { + if request.URL.Path != "/v1/files" { + t.Fatalf("path mismatch: got %s", request.URL.Path) + } + if request.URL.RawQuery != "" { + t.Fatalf("pagination controls must stay provider-local, got query %q", request.URL.RawQuery) + } + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"data":[{"id":"file-1","object":"file","filename":"one.txt","purpose":"general"},{"id":"file-2","object":"file","filename":"two.txt","purpose":"general"}]}`)) + })) + defer server.Close() + + provider := newTestGigaChatChatProvider(t, server.URL) + request := &schemas.BifrostFileListRequest{Provider: schemas.GigaChat, Limit: 1} + first, bifrostErr := provider.FileList(testBifrostContext(), []schemas.Key{testGigaChatAccessTokenKey("files-token")}, request) + if bifrostErr != nil { + t.Fatalf("first FileList returned error: %v", bifrostErr) + } + if len(first.Data) != 1 || first.Data[0].ID != "file-1" || !first.HasMore || first.After == nil { + t.Fatalf("unexpected first page: %#v", first) + } + + request.After = first.After + second, bifrostErr := provider.FileList(testBifrostContext(), []schemas.Key{testGigaChatAccessTokenKey("files-token")}, request) + if bifrostErr != nil { + t.Fatalf("second FileList returned error: %v", bifrostErr) + } + if len(second.Data) != 1 || second.Data[0].ID != "file-2" || second.HasMore || second.After != nil { + t.Fatalf("unexpected second page: %#v", second) + } +} + +func testGigaChatFileListRejectsUnsupportedOrder(t *testing.T) { + t.Parallel() + + provider, err := NewGigaChatProvider(&schemas.ProviderConfig{}, nil) + if err != nil { + t.Fatalf("NewGigaChatProvider returned error: %v", err) + } + order := "desc" + response, bifrostErr := provider.FileList(testBifrostContext(), []schemas.Key{testGigaChatAccessTokenKey("files-token")}, &schemas.BifrostFileListRequest{ + Provider: schemas.GigaChat, + Order: &order, + }) + if response != nil { + t.Fatalf("expected nil response, got %#v", response) + } + if bifrostErr == nil || !strings.Contains(bifrostErr.GetErrorString(), "does not support order sorting") { + t.Fatalf("unexpected error: %v", bifrostErr) + } +} + +func testGigaChatFileListPreservesUpstreamPurposeWhenFiltering(t *testing.T) { + t.Parallel() + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) { + if request.URL.Path != "/v1/files" { + t.Fatalf("path mismatch: got %s", request.URL.Path) + } + if request.Method != http.MethodGet { + t.Fatalf("method mismatch: got %s, want GET", request.Method) + } + + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"data":[{"id":"general-file","object":"file","bytes":10,"created_at":1780306293,"filename":"general.txt","purpose":"general"},{"id":"assistant-file","object":"file","bytes":20,"created_at":1780306294,"filename":"assistant.txt","purpose":"assistant"}]}`)) + })) + defer server.Close() + + provider := newTestGigaChatChatProvider(t, server.URL) + response, bifrostErr := provider.FileList(testBifrostContext(), []schemas.Key{testGigaChatAccessTokenKey("files-token")}, &schemas.BifrostFileListRequest{ + Provider: schemas.GigaChat, + Purpose: schemas.FilePurposeAssistants, + }) + if bifrostErr != nil { + t.Fatalf("FileList returned error: %v", bifrostErr) + } + if len(response.Data) != 1 { + t.Fatalf("file count mismatch: got %d files: %#v", len(response.Data), response.Data) + } + if response.Data[0].ID != "assistant-file" || response.Data[0].Purpose != schemas.FilePurposeAssistants { + t.Fatalf("unexpected filtered file: %#v", response.Data[0]) + } +} + +func testGigaChatFileListRetrieveDelete(t *testing.T) { + t.Parallel() + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) { + if got := request.Header.Get("Authorization"); got != "Bearer files-token" { + t.Fatalf("authorization header mismatch: got %q", got) + } + w.Header().Set("Content-Type", "application/json") + + switch request.URL.Path { + case "/v1/files": + if request.Method != http.MethodGet { + t.Fatalf("list method mismatch: got %s", request.Method) + } + _, _ = w.Write([]byte(`{"data":[{"id":"file-1","object":"file","bytes":10,"created_at":1780306293,"filename":"assistant.txt","purpose":"assistant"},{"id":"file-2","object":"file","bytes":20,"created_at":1780306294,"filename":"general.txt","purpose":"general"}]}`)) + case "/v1/files/file-1": + if request.Method != http.MethodGet { + t.Fatalf("retrieve method mismatch: got %s", request.Method) + } + _, _ = w.Write([]byte(`{"id":"file-1","object":"file","bytes":10,"created_at":1780306293,"filename":"assistant.txt","purpose":"assistant"}`)) + case "/v1/files/file-1/delete": + if request.Method != http.MethodPost { + t.Fatalf("delete method mismatch: got %s", request.Method) + } + _, _ = w.Write([]byte(`{"id":"file-1","deleted":true}`)) + default: + t.Fatalf("unexpected path: %s", request.URL.Path) + } + })) + defer server.Close() + + provider := newTestGigaChatChatProvider(t, server.URL) + key := testGigaChatAccessTokenKey("files-token") + ctx := testBifrostContext() + + listResponse, bifrostErr := provider.FileList(ctx, []schemas.Key{key}, &schemas.BifrostFileListRequest{Provider: schemas.GigaChat}) + if bifrostErr != nil { + t.Fatalf("FileList returned error: %v", bifrostErr) + } + if len(listResponse.Data) != 2 { + t.Fatalf("file count mismatch: got %d, want 2", len(listResponse.Data)) + } + if listResponse.Data[0].Purpose != schemas.FilePurposeAssistants { + t.Fatalf("assistant purpose mismatch: got %q", listResponse.Data[0].Purpose) + } + if listResponse.Data[1].Purpose != schemas.FilePurposeUserData { + t.Fatalf("general purpose mismatch: got %q", listResponse.Data[1].Purpose) + } + + retrieveResponse, bifrostErr := provider.FileRetrieve(ctx, []schemas.Key{key}, &schemas.BifrostFileRetrieveRequest{ + Provider: schemas.GigaChat, + FileID: "file-1", + }) + if bifrostErr != nil { + t.Fatalf("FileRetrieve returned error: %v", bifrostErr) + } + if retrieveResponse.ID != "file-1" || retrieveResponse.Purpose != schemas.FilePurposeAssistants { + t.Fatalf("unexpected retrieve response: %#v", retrieveResponse) + } + + deleteResponse, bifrostErr := provider.FileDelete(ctx, []schemas.Key{key}, &schemas.BifrostFileDeleteRequest{ + Provider: schemas.GigaChat, + FileID: "file-1", + }) + if bifrostErr != nil { + t.Fatalf("FileDelete returned error: %v", bifrostErr) + } + if deleteResponse.ID != "file-1" || !deleteResponse.Deleted || deleteResponse.Object != "file" { + t.Fatalf("unexpected delete response: %#v", deleteResponse) + } +} + +func testGigaChatFileContentRawBytes(t *testing.T) { + t.Parallel() + + wantContent := []byte{0xff, 0xd8, 0xff, 0xdb} + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) { + if request.URL.Path != "/v1/files/image-file/content" { + t.Fatalf("path mismatch: got %s", request.URL.Path) + } + if request.Method != http.MethodGet { + t.Fatalf("method mismatch: got %s, want GET", request.Method) + } + if got := request.Header.Get("Authorization"); got != "Bearer content-token" { + t.Fatalf("authorization header mismatch: got %q", got) + } + w.Header().Set("Content-Type", "image/jpeg") + _, _ = w.Write(wantContent) + })) + defer server.Close() + + provider := newTestGigaChatChatProvider(t, server.URL) + response, bifrostErr := provider.FileContent(testBifrostContext(), []schemas.Key{testGigaChatAccessTokenKey("content-token")}, &schemas.BifrostFileContentRequest{ + Provider: schemas.GigaChat, + FileID: "image-file", + }) + if bifrostErr != nil { + t.Fatalf("FileContent returned error: %v", bifrostErr) + } + if !bytes.Equal(response.Content, wantContent) { + t.Fatalf("content mismatch: got %v, want %v", response.Content, wantContent) + } + if response.ContentType != "image/jpeg" { + t.Fatalf("content type mismatch: got %q", response.ContentType) + } +} + +func testGigaChatFileContentBase64Wrapper(t *testing.T) { + t.Parallel() + + encoded := base64.StdEncoding.EncodeToString([]byte("decoded file content")) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) { + if request.URL.Path != "/v1/files/wrapped-file/content" { + t.Fatalf("path mismatch: got %s", request.URL.Path) + } + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"content":"` + encoded + `"}`)) + })) + defer server.Close() + + provider := newTestGigaChatChatProvider(t, server.URL) + response, bifrostErr := provider.FileContent(testBifrostContext(), []schemas.Key{testGigaChatAccessTokenKey("wrapper-token")}, &schemas.BifrostFileContentRequest{ + Provider: schemas.GigaChat, + FileID: "wrapped-file", + }) + if bifrostErr != nil { + t.Fatalf("FileContent returned error: %v", bifrostErr) + } + if string(response.Content) != "decoded file content" { + t.Fatalf("content mismatch: got %q", string(response.Content)) + } + if response.ContentType != "application/octet-stream" { + t.Fatalf("content type mismatch: got %q", response.ContentType) + } +} + +func testGigaChatFileUploadRefreshesTokenAfterUnauthorized(t *testing.T) { + t.Parallel() + + var tokenRequests atomic.Int32 + var uploadRequests atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) { + switch request.URL.Path { + case "/oauth": + tokenIndex := tokenRequests.Add(1) + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"access_token":"files-token-` + formatInt32(tokenIndex) + `","expires_at":1893456000}`)) + case "/v1/files": + uploadIndex := uploadRequests.Add(1) + wantAuthorization := "Bearer files-token-" + formatInt32(uploadIndex) + if got := request.Header.Get("Authorization"); got != wantAuthorization { + t.Fatalf("authorization header mismatch on request %d: got %q, want %q", uploadIndex, got, wantAuthorization) + } + w.Header().Set("Content-Type", "application/json") + if uploadIndex == 1 { + w.WriteHeader(http.StatusUnauthorized) + _, _ = w.Write([]byte(`{"status":401,"message":"expired token"}`)) + return + } + _, _ = w.Write([]byte(`{"id":"file-refreshed","object":"file","bytes":4,"created_at":1780306293,"filename":"file.txt","purpose":"general"}`)) + default: + t.Fatalf("unexpected path: %s", request.URL.Path) + } + })) + defer server.Close() + + provider := newTestGigaChatChatProvider(t, server.URL) + key := testGigaChatOAuthKey(server.URL+"/oauth", "", "test-credentials") + + response, bifrostErr := provider.FileUpload(testBifrostContext(), key, &schemas.BifrostFileUploadRequest{ + Provider: schemas.GigaChat, + File: []byte("test"), + Filename: "file.txt", + Purpose: schemas.FilePurposeUserData, + }) + if bifrostErr != nil { + t.Fatalf("FileUpload returned error: %v", bifrostErr) + } + if response.ID != "file-refreshed" { + t.Fatalf("unexpected response: %#v", response) + } + if tokenRequests.Load() != 2 { + t.Fatalf("token request count mismatch: got %d, want 2", tokenRequests.Load()) + } + if uploadRequests.Load() != 2 { + t.Fatalf("upload request count mismatch: got %d, want 2", uploadRequests.Load()) + } +} diff --git a/core/providers/gigachat/gigachat_comprehensive_test.go b/core/providers/gigachat/gigachat_comprehensive_test.go new file mode 100644 index 00000000000..e633ef16d9c --- /dev/null +++ b/core/providers/gigachat/gigachat_comprehensive_test.go @@ -0,0 +1,60 @@ +package gigachat_test + +import ( + "os" + "strings" + "testing" + + "github.com/maximhq/bifrost/core/internal/llmtests" +) + +func TestGigachat(t *testing.T) { + t.Parallel() + + validateGigaChatComprehensiveEnv(t) + if !hasGigaChatComprehensiveAuthEnv() { + t.Skip("Skipping GigaChat comprehensive tests because GIGACHAT_ACCESS_TOKEN, GIGACHAT_CREDENTIALS, or GIGACHAT_USER+GIGACHAT_PASSWORD+GIGACHAT_BASE_URL is not set") + } + + client, ctx, cancel, err := llmtests.SetupTest() + if err != nil { + t.Fatalf("Error initializing test setup: %v", err) + } + defer cancel() + defer client.Shutdown() + + t.Run("GigachatTests", func(t *testing.T) { + llmtests.RunAllComprehensiveTests(t, client, ctx, llmtests.GigaChatComprehensiveTestConfig()) + }) +} + +func hasGigaChatComprehensiveAuthEnv() bool { + if strings.TrimSpace(os.Getenv("GIGACHAT_ACCESS_TOKEN")) != "" { + return true + } + if strings.TrimSpace(os.Getenv("GIGACHAT_CREDENTIALS")) != "" { + return true + } + return strings.TrimSpace(os.Getenv("GIGACHAT_USER")) != "" && + strings.TrimSpace(os.Getenv("GIGACHAT_PASSWORD")) != "" && + strings.TrimSpace(os.Getenv("GIGACHAT_BASE_URL")) != "" +} + +func validateGigaChatComprehensiveEnv(t *testing.T) { + t.Helper() + + hasCertFile := strings.TrimSpace(os.Getenv("GIGACHAT_CERT_FILE")) != "" + hasKeyFile := strings.TrimSpace(os.Getenv("GIGACHAT_KEY_FILE")) != "" + if hasCertFile != hasKeyFile { + t.Fatal("GIGACHAT_CERT_FILE and GIGACHAT_KEY_FILE must be set together") + } + + hasUser := strings.TrimSpace(os.Getenv("GIGACHAT_USER")) != "" + hasPassword := strings.TrimSpace(os.Getenv("GIGACHAT_PASSWORD")) != "" + if hasUser != hasPassword { + t.Fatal("GIGACHAT_USER and GIGACHAT_PASSWORD must be set together") + } + if hasUser && strings.TrimSpace(os.Getenv("GIGACHAT_BASE_URL")) == "" { + t.Fatal("GIGACHAT_BASE_URL must be set when using GIGACHAT_USER and GIGACHAT_PASSWORD") + } +} diff --git a/core/providers/gigachat/gigachat_integration_test.go b/core/providers/gigachat/gigachat_integration_test.go new file mode 100644 index 00000000000..8e7bc22e9a0 --- /dev/null +++ b/core/providers/gigachat/gigachat_integration_test.go @@ -0,0 +1,751 @@ +package gigachat + +import ( + "bytes" + "context" + "errors" + "fmt" + "os" + "strings" + "testing" + "time" + + schemas "github.com/maximhq/bifrost/core/schemas" +) + +const ( + gigaChatIntegrationDefaultChatModel = "GigaChat-2" + gigaChatIntegrationDefaultEmbeddingModel = "Embeddings" + gigaChatIntegrationTimeout = 90 * time.Second +) + +type gigaChatIntegrationConfig struct { + baseURL string + authURL string + scope string + chatModel string + embeddingModel string + reasoningModel string + reasoningEffort string + batchModel string + webSearchModel string + hasAccessToken bool + hasOAuth bool + hasPassword bool + runBatch bool + runWebSearch bool + caBundleFile string + certFile string + keyFile string +} + +func TestGigaChatIntegration(t *testing.T) { + config := loadGigaChatIntegrationConfig(t) + provider := newGigaChatIntegrationProvider(t, config) + key := config.inferenceKey() + + t.Run("OAuthTokenFetch", func(t *testing.T) { + if !config.hasOAuth { + t.Skip("set GIGACHAT_CREDENTIALS and GIGACHAT_SCOPE to run OAuth token integration test") + } + ctx := newGigaChatIntegrationContext(t) + token, bifrostErr := provider.getOAuthAccessToken(ctx, config.oauthKey()) + if bifrostErr != nil { + failGigaChatIntegrationBifrostError(t, "OAuth token fetch", bifrostErr) + } + if strings.TrimSpace(token) == "" { + t.Fatal("OAuth token fetch returned an empty access token") + } + }) + + t.Run("PasswordTokenFetch", func(t *testing.T) { + if !config.hasPassword { + t.Skip("set GIGACHAT_USER, GIGACHAT_PASSWORD, and GIGACHAT_BASE_URL to run password token integration test") + } + ctx := newGigaChatIntegrationContext(t) + token, bifrostErr := provider.getPasswordAccessToken(ctx, config.passwordKey()) + if bifrostErr != nil { + failGigaChatIntegrationBifrostError(t, "password token fetch", bifrostErr) + } + if strings.TrimSpace(token) == "" { + t.Fatal("password token fetch returned an empty access token") + } + }) + + t.Run("ChatCompletion", func(t *testing.T) { + ctx := newGigaChatIntegrationContext(t) + response, bifrostErr := provider.ChatCompletion(ctx, key, gigaChatIntegrationChatRequest(config.chatModel)) + if bifrostErr != nil { + failGigaChatIntegrationBifrostError(t, "chat completion", bifrostErr) + } + if response == nil || len(response.Choices) == 0 { + t.Fatal("chat completion returned no choices") + } + if response.ExtraFields.Provider != schemas.GigaChat { + t.Fatalf("provider mismatch: got %q, want %q", response.ExtraFields.Provider, schemas.GigaChat) + } + }) + + t.Run("ChatCompletionStream", func(t *testing.T) { + ctx := newGigaChatIntegrationContext(t) + stream, bifrostErr := provider.ChatCompletionStream(ctx, testGigaChatPostHookRunner, nil, key, gigaChatIntegrationChatRequest(config.chatModel)) + if bifrostErr != nil { + failGigaChatIntegrationBifrostError(t, "chat completion stream", bifrostErr) + } + assertGigaChatIntegrationChatStream(t, stream) + }) + + t.Run("ListModels", func(t *testing.T) { + ctx := newGigaChatIntegrationContext(t) + response, bifrostErr := provider.ListModels(ctx, []schemas.Key{key}, &schemas.BifrostListModelsRequest{Provider: schemas.GigaChat}) + if bifrostErr != nil { + failGigaChatIntegrationBifrostError(t, "list models", bifrostErr) + } + if response == nil || len(response.Data) == 0 { + t.Fatal("list models returned no models") + } + }) + + t.Run("Files", func(t *testing.T) { + runGigaChatIntegrationFileLifecycle(t, provider, key) + }) + + t.Run("Embedding", func(t *testing.T) { + ctx := newGigaChatIntegrationContext(t) + response, bifrostErr := provider.Embedding(ctx, key, gigaChatIntegrationEmbeddingRequest(config.embeddingModel)) + if bifrostErr != nil { + failGigaChatIntegrationBifrostError(t, "embedding", bifrostErr) + } + if response == nil || len(response.Data) == 0 { + t.Fatal("embedding returned no vectors") + } + if response.ExtraFields.Provider != schemas.GigaChat { + t.Fatalf("provider mismatch: got %q, want %q", response.ExtraFields.Provider, schemas.GigaChat) + } + }) + + t.Run("Responses", func(t *testing.T) { + ctx := newGigaChatIntegrationContext(t) + response, bifrostErr := provider.Responses(ctx, key, gigaChatIntegrationResponsesRequest(config.chatModel)) + if bifrostErr != nil { + failGigaChatIntegrationBifrostError(t, "responses", bifrostErr) + } + if response == nil || len(response.Output) == 0 { + t.Fatal("responses returned no output") + } + if response.ExtraFields.Provider != schemas.GigaChat { + t.Fatalf("provider mismatch: got %q, want %q", response.ExtraFields.Provider, schemas.GigaChat) + } + }) + + t.Run("CountTokens", func(t *testing.T) { + ctx := newGigaChatIntegrationContext(t) + response, bifrostErr := provider.CountTokens(ctx, key, gigaChatIntegrationCountTokensRequest(config.chatModel)) + if bifrostErr != nil { + failGigaChatIntegrationBifrostError(t, "count tokens", bifrostErr) + } + if response == nil { + t.Fatal("count tokens returned nil response") + } + if response.InputTokens <= 0 { + t.Fatalf("count tokens returned invalid input token count: %#v", response) + } + if response.TotalTokens == nil || *response.TotalTokens < response.InputTokens { + t.Fatalf("count tokens returned invalid total tokens: %#v", response) + } + if len(response.Tokens) != 2 { + t.Fatalf("count tokens returned unexpected per-input counts: %#v", response.Tokens) + } + if response.ExtraFields.Provider != schemas.GigaChat { + t.Fatalf("provider mismatch: got %q, want %q", response.ExtraFields.Provider, schemas.GigaChat) + } + }) + + t.Run("ResponsesStream", func(t *testing.T) { + ctx := newGigaChatIntegrationContext(t) + stream, bifrostErr := provider.ResponsesStream(ctx, testGigaChatPostHookRunner, nil, key, gigaChatIntegrationResponsesRequest(config.chatModel)) + if bifrostErr != nil { + failGigaChatIntegrationBifrostError(t, "responses stream", bifrostErr) + } + assertGigaChatIntegrationResponsesStream(t, stream) + }) + + t.Run("ResponsesReasoning", func(t *testing.T) { + if config.reasoningModel == "" { + t.Skip("set GIGACHAT_REASONING_MODEL to run Responses reasoning integration test") + } + ctx := newGigaChatIntegrationContext(t) + response, bifrostErr := provider.Responses(ctx, key, gigaChatIntegrationReasoningRequest(config.reasoningModel, config.reasoningEffort)) + if bifrostErr != nil { + failGigaChatIntegrationBifrostError(t, "responses reasoning", bifrostErr) + } + if response == nil || len(response.Output) == 0 { + t.Fatal("responses reasoning returned no output") + } + if !gigaChatIntegrationHasReasoningOutput(response) { + t.Skip("GigaChat reasoning model did not return a reasoning output item for the smoke prompt") + } + }) + + t.Run("FunctionToolOptionalSchema", func(t *testing.T) { + ctx := newGigaChatIntegrationContext(t) + response, bifrostErr := provider.ChatCompletion(ctx, key, gigaChatIntegrationOptionalToolRequest(t, config.chatModel)) + if bifrostErr != nil { + failGigaChatIntegrationBifrostError(t, "function tool with optional schema", bifrostErr) + } + if !gigaChatIntegrationHasToolCall(response) { + t.Skip("GigaChat did not return a tool call for the forced tool request") + } + }) + + t.Run("BatchCreateRetrieve", func(t *testing.T) { + if !config.runBatch { + t.Skip("set GIGACHAT_ENABLE_BATCH_TEST=1 to run GigaChat batch integration test") + } + ctx := newGigaChatIntegrationContext(t) + response, bifrostErr := provider.BatchCreate(ctx, key, gigaChatIntegrationBatchCreateRequest(config.batchModel)) + if bifrostErr != nil { + failGigaChatIntegrationBifrostError(t, "batch create", bifrostErr) + } + if response == nil || strings.TrimSpace(response.ID) == "" { + t.Fatalf("batch create returned no batch ID: %#v", response) + } + + retrieveResponse, bifrostErr := provider.BatchRetrieve(ctx, []schemas.Key{key}, &schemas.BifrostBatchRetrieveRequest{ + Provider: schemas.GigaChat, + BatchID: response.ID, + }) + if bifrostErr != nil { + failGigaChatIntegrationBifrostError(t, "batch retrieve", bifrostErr) + } + if retrieveResponse == nil || retrieveResponse.ID != response.ID { + t.Fatalf("batch retrieve response mismatch: got %#v, want ID %q", retrieveResponse, response.ID) + } + }) + + t.Run("WebSearchBuiltIn", func(t *testing.T) { + if !config.runWebSearch { + t.Skip("set GIGACHAT_ENABLE_WEB_SEARCH_TEST=1 to run GigaChat web search integration test") + } + ctx := newGigaChatIntegrationContext(t) + response, bifrostErr := provider.Responses(ctx, key, gigaChatIntegrationWebSearchRequest(config.webSearchModel)) + if bifrostErr != nil { + failGigaChatIntegrationBifrostError(t, "web search built-in", bifrostErr) + } + if response == nil || len(response.Output) == 0 { + t.Fatal("web search built-in returned no output") + } + }) +} + +func TestGigaChatIntegrationConfigUsesAccessToken(t *testing.T) { + t.Setenv("GIGACHAT_ACCESS_TOKEN", "integration-test-access-token") + t.Setenv("GIGACHAT_CREDENTIALS", "") + t.Setenv("GIGACHAT_SCOPE", "") + t.Setenv("GIGACHAT_USER", "") + t.Setenv("GIGACHAT_PASSWORD", "") + t.Setenv("GIGACHAT_BASE_URL", "") + t.Setenv("GIGACHAT_CERT_FILE", "") + t.Setenv("GIGACHAT_KEY_FILE", "") + + config := loadGigaChatIntegrationConfig(t) + if !config.hasAccessToken { + t.Fatal("access-token-only integration config was not accepted") + } + if config.hasOAuth || config.hasPassword { + t.Fatalf("unexpected token flow flags: oauth=%t password=%t", config.hasOAuth, config.hasPassword) + } + + key := config.inferenceKey() + if key.Name != "gigachat-integration-access-token" { + t.Fatalf("inference key name mismatch: got %q", key.Name) + } + if key.GigaChatKeyConfig == nil || !key.GigaChatKeyConfig.AccessToken.IsSet() { + t.Fatalf("inference key did not use access-token auth: %#v", key.GigaChatKeyConfig) + } + if key.GigaChatKeyConfig.AccessToken.GetRawRef() != "env.GIGACHAT_ACCESS_TOKEN" { + t.Fatalf("access token env var mismatch: got %q", key.GigaChatKeyConfig.AccessToken.GetRawRef()) + } +} + +func loadGigaChatIntegrationConfig(t *testing.T) gigaChatIntegrationConfig { + t.Helper() + + baseURL := gigaChatIntegrationEnv("GIGACHAT_BASE_URL") + scope := gigaChatIntegrationEnv("GIGACHAT_SCOPE") + hasAccessToken := gigaChatIntegrationEnv("GIGACHAT_ACCESS_TOKEN") != "" + hasCredentials := gigaChatIntegrationEnv("GIGACHAT_CREDENTIALS") != "" + hasUser := gigaChatIntegrationEnv("GIGACHAT_USER") != "" + hasPasswordValue := gigaChatIntegrationEnv("GIGACHAT_PASSWORD") != "" + + config := gigaChatIntegrationConfig{ + baseURL: baseURL, + authURL: gigaChatIntegrationEnv("GIGACHAT_AUTH_URL"), + scope: scope, + chatModel: gigaChatIntegrationEnvWithDefault("GIGACHAT_CHAT_MODEL", gigaChatIntegrationDefaultChatModel), + embeddingModel: gigaChatIntegrationEnvWithDefault("GIGACHAT_EMBEDDING_MODEL", gigaChatIntegrationDefaultEmbeddingModel), + reasoningModel: gigaChatIntegrationEnv("GIGACHAT_REASONING_MODEL"), + reasoningEffort: gigaChatIntegrationEnvWithDefault("GIGACHAT_REASONING_EFFORT", "low"), + batchModel: gigaChatIntegrationEnvWithDefault("GIGACHAT_BATCH_MODEL", gigaChatIntegrationEnvWithDefault("GIGACHAT_CHAT_MODEL", gigaChatIntegrationDefaultChatModel)), + webSearchModel: gigaChatIntegrationEnvWithDefault("GIGACHAT_WEB_SEARCH_MODEL", gigaChatIntegrationEnvWithDefault("GIGACHAT_CHAT_MODEL", gigaChatIntegrationDefaultChatModel)), + hasAccessToken: hasAccessToken, + hasOAuth: hasCredentials && scope != "", + hasPassword: hasUser && hasPasswordValue && baseURL != "", + runBatch: gigaChatIntegrationEnvBool("GIGACHAT_ENABLE_BATCH_TEST"), + runWebSearch: gigaChatIntegrationEnvBool("GIGACHAT_ENABLE_WEB_SEARCH_TEST"), + caBundleFile: gigaChatIntegrationEnv("GIGACHAT_CA_BUNDLE_FILE"), + certFile: gigaChatIntegrationEnv("GIGACHAT_CERT_FILE"), + keyFile: gigaChatIntegrationEnv("GIGACHAT_KEY_FILE"), + } + + if !config.hasAccessToken && !config.hasOAuth && !config.hasPassword { + t.Skip("set either GIGACHAT_ACCESS_TOKEN, GIGACHAT_CREDENTIALS+GIGACHAT_SCOPE, or GIGACHAT_USER+GIGACHAT_PASSWORD+GIGACHAT_BASE_URL to run GigaChat integration tests") + } + if (config.certFile == "") != (config.keyFile == "") { + t.Fatal("GIGACHAT_CERT_FILE and GIGACHAT_KEY_FILE must be set together") + } + + return config +} + +func gigaChatIntegrationEnv(name string) string { + return strings.TrimSpace(os.Getenv(name)) +} + +func gigaChatIntegrationEnvWithDefault(name string, defaultValue string) string { + value := gigaChatIntegrationEnv(name) + if value == "" { + return defaultValue + } + return value +} + +func gigaChatIntegrationEnvBool(name string) bool { + switch strings.ToLower(gigaChatIntegrationEnv(name)) { + case "1", "true", "yes", "y", "on": + return true + default: + return false + } +} + +func newGigaChatIntegrationProvider(t *testing.T, config gigaChatIntegrationConfig) *GigaChatProvider { + t.Helper() + + providerConfig := &schemas.ProviderConfig{} + if config.baseURL != "" { + providerConfig.NetworkConfig.BaseURL = config.baseURL + } + provider, err := NewGigaChatProvider(providerConfig, nil) + if err != nil { + failGigaChatIntegrationError(t, "new provider", err) + } + return provider +} + +func (config gigaChatIntegrationConfig) inferenceKey() schemas.Key { + if config.hasAccessToken { + return config.accessTokenKey() + } + if config.hasOAuth { + return config.oauthKey() + } + return config.passwordKey() +} + +func (config gigaChatIntegrationConfig) accessTokenKey() schemas.Key { + keyConfig := config.keyConfig() + keyConfig.AccessToken = schemas.NewSecretVar("env.GIGACHAT_ACCESS_TOKEN") + return schemas.Key{ + Name: "gigachat-integration-access-token", + Models: schemas.WhiteList{"*"}, + GigaChatKeyConfig: keyConfig, + } +} + +func (config gigaChatIntegrationConfig) oauthKey() schemas.Key { + keyConfig := config.keyConfig() + keyConfig.Credentials = schemas.NewSecretVar("env.GIGACHAT_CREDENTIALS") + keyConfig.Scope = config.scope + return schemas.Key{ + Name: "gigachat-integration-oauth", + Models: schemas.WhiteList{"*"}, + GigaChatKeyConfig: keyConfig, + } +} + +func (config gigaChatIntegrationConfig) passwordKey() schemas.Key { + keyConfig := config.keyConfig() + keyConfig.User = schemas.NewSecretVar("env.GIGACHAT_USER") + keyConfig.Password = schemas.NewSecretVar("env.GIGACHAT_PASSWORD") + return schemas.Key{ + Name: "gigachat-integration-password", + Models: schemas.WhiteList{"*"}, + GigaChatKeyConfig: keyConfig, + } +} + +func (config gigaChatIntegrationConfig) keyConfig() *schemas.GigaChatKeyConfig { + return &schemas.GigaChatKeyConfig{ + AuthURL: config.authURL, + BaseURL: config.baseURL, + CertFile: config.certFile, + KeyFile: config.keyFile, + CABundleFile: config.caBundleFile, + } +} + +func newGigaChatIntegrationContext(t *testing.T) *schemas.BifrostContext { + t.Helper() + + ctx, cancel := schemas.NewBifrostContextWithTimeout(context.Background(), gigaChatIntegrationTimeout) + t.Cleanup(cancel) + return ctx +} + +func gigaChatIntegrationChatRequest(model string) *schemas.BifrostChatRequest { + maxTokens := 64 + text := "Reply with one short sentence." + return &schemas.BifrostChatRequest{ + Model: model, + Input: []schemas.ChatMessage{{ + Role: schemas.ChatMessageRoleUser, + Content: &schemas.ChatMessageContent{ContentStr: &text}, + }}, + Params: &schemas.ChatParameters{ + MaxCompletionTokens: &maxTokens, + }, + } +} + +func gigaChatIntegrationEmbeddingRequest(model string) *schemas.BifrostEmbeddingRequest { + return &schemas.BifrostEmbeddingRequest{ + Model: model, + Input: &schemas.EmbeddingInput{Text: schemas.Ptr("integration test")}, + } +} + +func gigaChatIntegrationResponsesRequest(model string) *schemas.BifrostResponsesRequest { + maxTokens := 64 + return &schemas.BifrostResponsesRequest{ + Model: model, + Input: []schemas.ResponsesMessage{{ + Role: schemas.Ptr(schemas.ResponsesInputMessageRoleUser), + Content: &schemas.ResponsesMessageContent{ContentStr: schemas.Ptr("Reply with one short sentence.")}, + }}, + Params: &schemas.ResponsesParameters{ + MaxOutputTokens: &maxTokens, + }, + } +} + +func gigaChatIntegrationCountTokensRequest(model string) *schemas.BifrostResponsesRequest { + return &schemas.BifrostResponsesRequest{ + Model: model, + Input: []schemas.ResponsesMessage{ + { + Role: schemas.Ptr(schemas.ResponsesInputMessageRoleUser), + Content: &schemas.ResponsesMessageContent{ContentStr: schemas.Ptr("Привет, как дела?")}, + }, + { + Role: schemas.Ptr(schemas.ResponsesInputMessageRoleUser), + Content: &schemas.ResponsesMessageContent{ContentStr: schemas.Ptr("Hello, how are you?")}, + }, + }, + } +} + +func gigaChatIntegrationReasoningRequest(model string, effort string) *schemas.BifrostResponsesRequest { + maxTokens := 128 + return &schemas.BifrostResponsesRequest{ + Model: model, + Input: []schemas.ResponsesMessage{{ + Role: schemas.Ptr(schemas.ResponsesInputMessageRoleUser), + Content: &schemas.ResponsesMessageContent{ContentStr: schemas.Ptr("Think briefly, then answer with exactly: four.")}, + }}, + Params: &schemas.ResponsesParameters{ + MaxOutputTokens: &maxTokens, + Reasoning: &schemas.ResponsesParametersReasoning{ + Effort: &effort, + }, + }, + } +} + +func gigaChatIntegrationOptionalToolRequest(t *testing.T, model string) *schemas.BifrostChatRequest { + t.Helper() + + request := testGigaChatChatToolRequest(t, "get_weather") + request.Model = model + request.Input[0].Content = &schemas.ChatMessageContent{ContentStr: schemas.Ptr("Use the get_weather function for Moscow. Leave units empty if unknown.")} + request.Params.Tools[0].Function.Parameters = mustGigaChatToolParameters(t, `{ + "type": "object", + "properties": { + "city": {"anyOf": [{"type": "string"}, {"type": "null"}]}, + "units": {"type": ["string", "null"], "nullable": true} + }, + "required": ["city"] + }`) + request.Params.ToolChoice = &schemas.ChatToolChoice{ + ChatToolChoiceStruct: &schemas.ChatToolChoiceStruct{ + Type: schemas.ChatToolChoiceTypeFunction, + Function: &schemas.ChatToolChoiceFunction{ + Name: "get_weather", + }, + }, + } + return request +} + +func gigaChatIntegrationBatchCreateRequest(model string) *schemas.BifrostBatchCreateRequest { + return &schemas.BifrostBatchCreateRequest{ + Provider: schemas.GigaChat, + Endpoint: schemas.BatchEndpointChatCompletions, + CompletionWindow: "24h", + Requests: []schemas.BatchRequestItem{{ + CustomID: fmt.Sprintf("bifrost-gigachat-integration-%d", time.Now().UnixNano()), + Method: "POST", + URL: string(schemas.BatchEndpointChatCompletions), + Body: map[string]interface{}{ + "model": model, + "messages": []map[string]interface{}{{ + "role": "user", + "content": "Reply with one short sentence.", + }}, + "max_tokens": 32, + }, + }}, + } +} + +func gigaChatIntegrationWebSearchRequest(model string) *schemas.BifrostResponsesRequest { + maxTokens := 128 + searchContextSize := "low" + return &schemas.BifrostResponsesRequest{ + Model: model, + Input: []schemas.ResponsesMessage{{ + Role: schemas.Ptr(schemas.ResponsesInputMessageRoleUser), + Content: &schemas.ResponsesMessageContent{ContentStr: schemas.Ptr("Use web search and answer in one sentence: what is the official GigaChat developer documentation site?")}, + }}, + Params: &schemas.ResponsesParameters{ + MaxOutputTokens: &maxTokens, + Tools: []schemas.ResponsesTool{{ + Type: schemas.ResponsesToolTypeWebSearchPreview, + ResponsesToolWebSearchPreview: &schemas.ResponsesToolWebSearchPreview{ + SearchContextSize: &searchContextSize, + }, + }}, + ToolChoice: &schemas.ResponsesToolChoice{ResponsesToolChoiceStruct: &schemas.ResponsesToolChoiceStruct{ + Type: schemas.ResponsesToolChoiceTypeWebSearchPreview, + }}, + }, + } +} + +func runGigaChatIntegrationFileLifecycle(t *testing.T, provider *GigaChatProvider, key schemas.Key) { + t.Helper() + + ctx := newGigaChatIntegrationContext(t) + content := []byte("bifrost gigachat integration file\n") + contentType := "text/plain" + uploadResponse, bifrostErr := provider.FileUpload(ctx, key, &schemas.BifrostFileUploadRequest{ + Provider: schemas.GigaChat, + File: content, + Filename: fmt.Sprintf("bifrost-gigachat-integration-%d.txt", time.Now().UnixNano()), + Purpose: schemas.FilePurposeUserData, + ContentType: &contentType, + }) + if bifrostErr != nil { + failGigaChatIntegrationBifrostError(t, "file upload", bifrostErr) + } + if uploadResponse == nil || strings.TrimSpace(uploadResponse.ID) == "" { + t.Fatalf("file upload returned no file ID: %#v", uploadResponse) + } + + fileID := uploadResponse.ID + deleted := false + defer func() { + if deleted { + return + } + if _, cleanupErr := provider.FileDelete(ctx, []schemas.Key{key}, &schemas.BifrostFileDeleteRequest{ + Provider: schemas.GigaChat, + FileID: fileID, + }); cleanupErr != nil { + t.Logf("cleanup file delete failed: %s", redactGigaChatIntegrationSecrets(cleanupErr.String())) + } + }() + + listResponse, bifrostErr := provider.FileList(ctx, []schemas.Key{key}, &schemas.BifrostFileListRequest{Provider: schemas.GigaChat}) + if bifrostErr != nil { + failGigaChatIntegrationBifrostError(t, "file list", bifrostErr) + } + if listResponse == nil || !gigaChatIntegrationFileListContains(listResponse.Data, fileID) { + t.Fatalf("file list did not include uploaded file %q", fileID) + } + + retrieveResponse, bifrostErr := provider.FileRetrieve(ctx, []schemas.Key{key}, &schemas.BifrostFileRetrieveRequest{ + Provider: schemas.GigaChat, + FileID: fileID, + }) + if bifrostErr != nil { + failGigaChatIntegrationBifrostError(t, "file retrieve", bifrostErr) + } + if retrieveResponse == nil || retrieveResponse.ID != fileID { + t.Fatalf("file retrieve response mismatch: got %#v, want ID %q", retrieveResponse, fileID) + } + + contentResponse, bifrostErr := provider.FileContent(ctx, []schemas.Key{key}, &schemas.BifrostFileContentRequest{ + Provider: schemas.GigaChat, + FileID: fileID, + }) + if bifrostErr != nil { + failGigaChatIntegrationBifrostError(t, "file content", bifrostErr) + } + if contentResponse == nil || !bytes.Equal(contentResponse.Content, content) { + t.Fatalf("file content mismatch: got %q, want %q", string(contentResponse.Content), string(content)) + } + + deleteResponse, bifrostErr := provider.FileDelete(ctx, []schemas.Key{key}, &schemas.BifrostFileDeleteRequest{ + Provider: schemas.GigaChat, + FileID: fileID, + }) + if bifrostErr != nil { + failGigaChatIntegrationBifrostError(t, "file delete", bifrostErr) + } + if deleteResponse == nil || deleteResponse.ID != fileID || !deleteResponse.Deleted { + t.Fatalf("file delete response mismatch: %#v", deleteResponse) + } + deleted = true +} + +func gigaChatIntegrationFileListContains(files []schemas.FileObject, fileID string) bool { + for _, file := range files { + if file.ID == fileID { + return true + } + } + return false +} + +func assertGigaChatIntegrationChatStream(t *testing.T, stream chan *schemas.BifrostStreamChunk) { + t.Helper() + + if stream == nil { + t.Fatal("chat completion stream is nil") + } + + receivedResponse := false + for chunk := range stream { + if chunk == nil { + continue + } + if chunk.BifrostError != nil { + failGigaChatIntegrationBifrostError(t, "chat completion stream chunk", chunk.BifrostError) + } + if chunk.BifrostChatResponse != nil { + receivedResponse = true + if chunk.BifrostChatResponse.ExtraFields.Provider != schemas.GigaChat { + t.Fatalf("provider mismatch: got %q, want %q", chunk.BifrostChatResponse.ExtraFields.Provider, schemas.GigaChat) + } + } + } + if !receivedResponse { + t.Fatal("chat completion stream returned no response chunks") + } +} + +func assertGigaChatIntegrationResponsesStream(t *testing.T, stream chan *schemas.BifrostStreamChunk) { + t.Helper() + + if stream == nil { + t.Fatal("responses stream is nil") + } + + receivedCompleted := false + for chunk := range stream { + if chunk == nil { + continue + } + if chunk.BifrostError != nil { + failGigaChatIntegrationBifrostError(t, "responses stream chunk", chunk.BifrostError) + } + if chunk.BifrostResponsesStreamResponse == nil { + continue + } + if chunk.BifrostResponsesStreamResponse.Response != nil && + chunk.BifrostResponsesStreamResponse.Response.ExtraFields.Provider != schemas.GigaChat { + t.Fatalf("provider mismatch: got %q, want %q", chunk.BifrostResponsesStreamResponse.Response.ExtraFields.Provider, schemas.GigaChat) + } + if chunk.BifrostResponsesStreamResponse.Type == schemas.ResponsesStreamResponseTypeCompleted { + receivedCompleted = true + } + } + if !receivedCompleted { + t.Fatal("responses stream did not emit response.completed") + } +} + +func gigaChatIntegrationHasReasoningOutput(response *schemas.BifrostResponsesResponse) bool { + if response == nil { + return false + } + for _, output := range response.Output { + if output.Type != nil && *output.Type == schemas.ResponsesMessageTypeReasoning { + return true + } + if output.ResponsesReasoning != nil && len(output.ResponsesReasoning.Summary) > 0 { + return true + } + } + return false +} + +func gigaChatIntegrationHasToolCall(response *schemas.BifrostChatResponse) bool { + if response == nil { + return false + } + for _, choice := range response.Choices { + if choice.ChatNonStreamResponseChoice == nil || choice.ChatNonStreamResponseChoice.Message == nil { + continue + } + message := choice.ChatNonStreamResponseChoice.Message + if message.ChatAssistantMessage != nil && len(message.ChatAssistantMessage.ToolCalls) > 0 { + return true + } + } + return false +} + +func failGigaChatIntegrationBifrostError(t *testing.T, operation string, bifrostErr *schemas.BifrostError) { + t.Helper() + + if bifrostErr == nil { + t.Fatalf("%s failed", operation) + } + failGigaChatIntegrationError(t, operation, errors.New(bifrostErr.String())) +} + +func failGigaChatIntegrationError(t *testing.T, operation string, err error) { + t.Helper() + + message := fmt.Sprintf("%s failed: %v", operation, err) + t.Fatal(redactGigaChatIntegrationSecrets(message)) +} + +func redactGigaChatIntegrationSecrets(message string) string { + for _, envName := range []string{ + "GIGACHAT_ACCESS_TOKEN", + "GIGACHAT_CREDENTIALS", + "GIGACHAT_USER", + "GIGACHAT_PASSWORD", + "GIGACHAT_CERT_FILE", + "GIGACHAT_KEY_FILE", + "GIGACHAT_CA_BUNDLE_FILE", + } { + if value := os.Getenv(envName); strings.TrimSpace(value) != "" { + message = strings.ReplaceAll(message, value, "") + } + } + return message +} diff --git a/core/providers/gigachat/gigachat_test.go b/core/providers/gigachat/gigachat_test.go new file mode 100644 index 00000000000..7a5abb7adb9 --- /dev/null +++ b/core/providers/gigachat/gigachat_test.go @@ -0,0 +1,533 @@ +package gigachat + +import ( + "context" + "crypto/rand" + "crypto/rsa" + "crypto/x509" + "crypto/x509/pkix" + "encoding/pem" + "io" + "math/big" + "os" + "strings" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/maximhq/bifrost/core/schemas" + "github.com/valyala/fasthttp" +) + +func TestGigachat(t *testing.T) { + t.Parallel() + + t.Run("NewProvider", testNewGigaChatProvider) + t.Run("TrimBaseURL", testNewGigaChatProviderTrimsBaseURL) + t.Run("UnsupportedOperation", testGigaChatProviderUnsupportedOperation) + t.Run("ChatCompletion", testGigaChatChatCompletion) + t.Run("ListModels", testGigaChatListModels) + t.Run("Embedding", testGigaChatEmbedding) + t.Run("ResponsesRequestConversion", testGigaChatResponsesRequestConversion) + t.Run("Responses", testGigaChatResponses) + t.Run("ResponsesStream", testGigaChatResponsesStream) + t.Run("Tools", testGigaChatTools) + t.Run("Errors", testGigaChatErrors) + t.Run("BuildsTLSClientWithCABundle", testGigaChatBuildsTLSClientWithCABundle) + t.Run("ReusesTLSClientWithCABundleUntilProviderReload", testGigaChatReusesTLSClientWithCABundleUntilProviderReload) + t.Run("CachesTLSClientConcurrently", testGigaChatCachesTLSClientConcurrently) + t.Run("BuildsTLSClientWithCertificate", testGigaChatBuildsTLSClientWithCertificate) + t.Run("ReusesTLSClientWithCertificateUntilProviderReload", testGigaChatReusesTLSClientWithCertificateUntilProviderReload) + t.Run("RejectsMissingCertificatePair", testGigaChatRejectsMissingCertificatePair) + t.Run("PassthroughEarlyCloseFinalizesOnce", testGigaChatPassthroughEarlyCloseFinalizesOnce) + t.Run("AttachmentCacheLifecycle", testGigaChatAttachmentCacheLifecycle) +} + +type gigaChatCountingReadCloser struct { + io.Reader + closeCalls atomic.Int32 +} + +func (reader *gigaChatCountingReadCloser) Close() error { + reader.closeCalls.Add(1) + return nil +} + +func testGigaChatPassthroughEarlyCloseFinalizesOnce(t *testing.T) { + t.Parallel() + + ctx := testBifrostContext() + underlying := &gigaChatCountingReadCloser{Reader: strings.NewReader("incomplete")} + var finalizerCalls atomic.Int32 + reader := &gigaChatPassthroughReadCloser{ + ReadCloser: underlying, + ctx: ctx, + postHookSpanFinalizer: func(context.Context) { + finalizerCalls.Add(1) + }, + } + + buffer := make([]byte, 1) + if _, err := reader.Read(buffer); err != nil { + t.Fatalf("failed to read passthrough prefix: %v", err) + } + + closeErrors := make(chan error, 2) + for range 2 { + go func() { + closeErrors <- reader.Close() + }() + } + for range 2 { + if err := <-closeErrors; err != nil { + t.Fatalf("passthrough close failed: %v", err) + } + } + + if got := underlying.closeCalls.Load(); got != 1 { + t.Fatalf("underlying close calls mismatch: got %d, want 1", got) + } + if got := finalizerCalls.Load(); got != 1 { + t.Fatalf("finalizer calls mismatch: got %d, want 1", got) + } + if ended, _ := ctx.Value(schemas.BifrostContextKeyStreamEndIndicator).(bool); ended { + t.Fatal("early passthrough close must not mark the stream complete") + } +} + +func testGigaChatAttachmentCacheLifecycle(t *testing.T) { + t.Parallel() + + manager := newGigaChatAttachmentCacheManager() + defer func() { + manager.mu.Lock() + defer manager.mu.Unlock() + if manager.sweepTimer != nil { + manager.sweepTimer.Stop() + manager.sweepTimer = nil + } + }() + + provider := &GigaChatProvider{attachmentCache: manager} + ctx, cancel := schemas.NewBifrostContextWithCancel(context.Background()) + request := &schemas.BifrostChatRequest{} + fileID := "uploaded-file" + replacement := schemas.ChatContentBlock{ + Type: schemas.ChatContentBlockTypeFile, + File: &schemas.ChatInputFile{FileID: &fileID}, + } + provider.setCachedGigaChatChatAttachment(ctx, schemas.Key{}, request, 0, 0, replacement) + + cacheID, ok := ctx.Value(gigaChatAttachmentCacheKey).(string) + if !ok || cacheID == "" { + t.Fatalf("context must store only a cache ID, got %#v", ctx.Value(gigaChatAttachmentCacheKey)) + } + manager.mu.Lock() + entry := manager.entries[cacheID] + manager.mu.Unlock() + if entry == nil { + t.Fatal("provider-owned attachment cache entry is missing") + } + childCtx := schemas.NewBifrostContext(ctx, schemas.NoDeadline) + childCacheID, _ := childCtx.Value(gigaChatAttachmentCacheKey).(string) + if childCacheID != cacheID { + t.Fatalf("derived context did not inherit cache ID: got %q, want %q", childCacheID, cacheID) + } + if cached, found := provider.getCachedGigaChatChatAttachment(childCtx, schemas.Key{}, request, 0, 0); !found || cached.File == nil || cached.File.FileID == nil || *cached.File.FileID != fileID { + t.Fatalf("derived context did not reuse cached attachment: %#v, found=%v", cached, found) + } + manager.mu.Lock() + entryCount := len(manager.entries) + manager.mu.Unlock() + if entryCount != 1 { + t.Fatalf("derived context created or evicted a cache bucket: got %d entries, want 1", entryCount) + } + + entry.cache.mu.Lock() + for key, attachment := range entry.cache.chat { + attachment.expiresAt = time.Time{} + entry.cache.chat[key] = attachment + } + entry.cache.mu.Unlock() + if cached, found := provider.getCachedGigaChatChatAttachment(ctx, schemas.Key{}, request, 0, 0); found { + t.Fatalf("expired attachment remained reusable: %#v", cached) + } + entry.cache.mu.Lock() + remainingAttachments := len(entry.cache.chat) + entry.cache.mu.Unlock() + if remainingAttachments != 0 { + t.Fatalf("expired attachment records were not pruned: %d remain", remainingAttachments) + } + + _, writer := manager.cacheForWrite(ctx) + if writer == nil { + t.Fatal("failed to register in-flight attachment cache writer") + } + manager.mu.Lock() + if manager.sweepTimer != nil { + manager.sweepTimer.Stop() + manager.sweepTimer = nil + } + manager.mu.Unlock() + manager.sweep() + manager.mu.Lock() + entriesDuringWrite := len(manager.entries) + manager.mu.Unlock() + if entriesDuringWrite != 1 { + t.Fatalf("sweep evicted a cache with an in-flight writer: got %d entries, want 1", entriesDuringWrite) + } + + manager.finishWrite(writer) + manager.mu.Lock() + if manager.sweepTimer != nil { + manager.sweepTimer.Stop() + manager.sweepTimer = nil + } + manager.mu.Unlock() + manager.sweep() + manager.mu.Lock() + remainingEntries := len(manager.entries) + manager.mu.Unlock() + if remainingEntries != 0 { + t.Fatalf("empty context cache was not pruned: %d entries remain", remainingEntries) + } + cancel() +} + +func testNewGigaChatProvider(t *testing.T) { + t.Parallel() + + provider, err := NewGigaChatProvider(&schemas.ProviderConfig{}, nil) + if err != nil { + t.Fatalf("NewGigaChatProvider returned error: %v", err) + } + if provider.GetProviderKey() != schemas.GigaChat { + t.Fatalf("provider key mismatch: got %q, want %q", provider.GetProviderKey(), schemas.GigaChat) + } + if provider.networkConfig.BaseURL != gigaChatDefaultBaseURL { + t.Fatalf("base URL mismatch: got %q, want %q", provider.networkConfig.BaseURL, gigaChatDefaultBaseURL) + } + if provider.client == nil { + t.Fatal("client is nil") + } + if provider.streamingClient == nil { + t.Fatal("streaming client is nil") + } +} + +func testNewGigaChatProviderTrimsBaseURL(t *testing.T) { + t.Parallel() + + provider, err := NewGigaChatProvider(&schemas.ProviderConfig{ + NetworkConfig: schemas.NetworkConfig{ + BaseURL: "https://api.giga.chat/v1/", + }, + }, nil) + if err != nil { + t.Fatalf("NewGigaChatProvider returned error: %v", err) + } + + if provider.networkConfig.BaseURL != "https://api.giga.chat/v1" { + t.Fatalf("base URL mismatch: got %q", provider.networkConfig.BaseURL) + } +} + +func testGigaChatProviderUnsupportedOperation(t *testing.T) { + t.Parallel() + + provider, err := NewGigaChatProvider(&schemas.ProviderConfig{}, nil) + if err != nil { + t.Fatalf("NewGigaChatProvider returned error: %v", err) + } + + response, bifrostErr := provider.TextCompletion(nil, schemas.Key{}, &schemas.BifrostTextCompletionRequest{}) + if response != nil { + t.Fatalf("expected nil response, got %#v", response) + } + if bifrostErr == nil { + t.Fatal("expected unsupported operation error, got nil") + } + if bifrostErr.Error == nil || bifrostErr.Error.Code == nil || *bifrostErr.Error.Code != "unsupported_operation" { + t.Fatalf("unexpected error code: %#v", bifrostErr.Error) + } + if bifrostErr.ExtraFields.Provider != schemas.GigaChat { + t.Fatalf("provider mismatch: got %q, want %q", bifrostErr.ExtraFields.Provider, schemas.GigaChat) + } + if bifrostErr.ExtraFields.RequestType != schemas.TextCompletionRequest { + t.Fatalf("request type mismatch: got %q, want %q", bifrostErr.ExtraFields.RequestType, schemas.TextCompletionRequest) + } + if !strings.Contains(bifrostErr.Error.Message, "gigachat provider") { + t.Fatalf("unexpected error message: %q", bifrostErr.Error.Message) + } +} + +func testGigaChatBuildsTLSClientWithCABundle(t *testing.T) { + t.Parallel() + + certPEM, _ := generateGigaChatTestCertificate(t) + caBundleFile := writeGigaChatTestFile(t, "ca.pem", certPEM) + + provider, err := NewGigaChatProvider(&schemas.ProviderConfig{}, nil) + if err != nil { + t.Fatalf("NewGigaChatProvider returned error: %v", err) + } + + client, err := buildGigaChatTLSClient(provider.client, &schemas.GigaChatKeyConfig{CABundleFile: caBundleFile}) + if err != nil { + t.Fatalf("buildGigaChatTLSClient returned error: %v", err) + } + if client == provider.client { + t.Fatal("expected a cloned client when TLS material is configured") + } + if client.TLSConfig == nil || client.TLSConfig.RootCAs == nil { + t.Fatalf("expected RootCAs to be configured, got %#v", client.TLSConfig) + } + if provider.client.TLSConfig != nil && provider.client.TLSConfig.RootCAs != nil { + t.Fatal("base client TLS config was mutated") + } + if client.MaxConnsPerHost != provider.client.MaxConnsPerHost { + t.Fatalf("MaxConnsPerHost mismatch: got %d, want %d", client.MaxConnsPerHost, provider.client.MaxConnsPerHost) + } + if client.ConnPoolStrategy != fasthttp.FIFO { + t.Fatalf("ConnPoolStrategy mismatch: got %v", client.ConnPoolStrategy) + } +} + +func testGigaChatReusesTLSClientWithCABundleUntilProviderReload(t *testing.T) { + t.Parallel() + + certPEM1, _ := generateGigaChatTestCertificate(t) + certPEM2, _ := generateGigaChatTestCertificate(t) + caBundleFile := writeGigaChatTestFile(t, "ca.pem", certPEM1) + + provider, err := NewGigaChatProvider(&schemas.ProviderConfig{}, nil) + if err != nil { + t.Fatalf("NewGigaChatProvider returned error: %v", err) + } + + keyConfig := &schemas.GigaChatKeyConfig{CABundleFile: caBundleFile} + client, err := provider.getGigaChatTLSClient(provider.client, gigaChatTLSClientCacheDefault, keyConfig) + if err != nil { + t.Fatalf("getGigaChatTLSClient returned error: %v", err) + } + if client == provider.client { + t.Fatal("expected a cloned client when TLS material is configured") + } + + if err := os.WriteFile(caBundleFile, certPEM2, 0o600); err != nil { + t.Fatalf("failed to rotate CA bundle file: %v", err) + } + + reusedClient, err := provider.getGigaChatTLSClient(provider.client, gigaChatTLSClientCacheDefault, keyConfig) + if err != nil { + t.Fatalf("getGigaChatTLSClient returned error after CA rotation: %v", err) + } + if reusedClient != client { + t.Fatal("expected cached TLS client to be reused after CA bundle rotation") + } + + if err := os.Remove(caBundleFile); err != nil { + t.Fatalf("failed to remove CA bundle file: %v", err) + } + reusedClient, err = provider.getGigaChatTLSClient(provider.client, gigaChatTLSClientCacheDefault, keyConfig) + if err != nil { + t.Fatalf("cached TLS client lookup read removed CA bundle file: %v", err) + } + if reusedClient != client { + t.Fatal("expected cached TLS client to be reused after CA bundle removal") + } + + reloadedProvider, err := NewGigaChatProvider(&schemas.ProviderConfig{}, nil) + if err != nil { + t.Fatalf("NewGigaChatProvider returned error: %v", err) + } + if _, err := reloadedProvider.getGigaChatTLSClient(reloadedProvider.client, gigaChatTLSClientCacheDefault, keyConfig); err == nil { + t.Fatal("expected provider reload to validate missing CA bundle file") + } else if !strings.Contains(err.Error(), "ca_bundle_file") { + t.Fatalf("unexpected missing CA error: %v", err) + } +} + +func testGigaChatCachesTLSClientConcurrently(t *testing.T) { + t.Parallel() + + certPEM, _ := generateGigaChatTestCertificate(t) + caBundleFile := writeGigaChatTestFile(t, "ca.pem", certPEM) + provider, err := NewGigaChatProvider(&schemas.ProviderConfig{}, nil) + if err != nil { + t.Fatalf("NewGigaChatProvider returned error: %v", err) + } + + const workerCount = 32 + keyConfig := &schemas.GigaChatKeyConfig{CABundleFile: caBundleFile} + clients := make([]*fasthttp.Client, workerCount) + errs := make([]error, workerCount) + start := make(chan struct{}) + var workers sync.WaitGroup + for i := range workerCount { + workers.Add(1) + go func() { + defer workers.Done() + <-start + clients[i], errs[i] = provider.getGigaChatTLSClient(provider.client, gigaChatTLSClientCacheDefault, keyConfig) + }() + } + close(start) + workers.Wait() + + for i, err := range errs { + if err != nil { + t.Fatalf("getGigaChatTLSClient call %d returned error: %v", i, err) + } + } + cachedClient := clients[0] + for _, client := range clients[1:] { + if client != cachedClient { + t.Fatal("concurrent cache misses returned different TLS clients") + } + } +} + +func testGigaChatBuildsTLSClientWithCertificate(t *testing.T) { + t.Parallel() + + certPEM, keyPEM := generateGigaChatTestCertificate(t) + certFile := writeGigaChatTestFile(t, "client.pem", certPEM) + keyFile := writeGigaChatTestFile(t, "client.key", keyPEM) + + provider, err := NewGigaChatProvider(&schemas.ProviderConfig{}, nil) + if err != nil { + t.Fatalf("NewGigaChatProvider returned error: %v", err) + } + + client, err := buildGigaChatTLSClient(provider.client, &schemas.GigaChatKeyConfig{ + CertFile: certFile, + KeyFile: keyFile, + }) + if err != nil { + t.Fatalf("buildGigaChatTLSClient returned error: %v", err) + } + if client.TLSConfig == nil || len(client.TLSConfig.Certificates) != 1 { + t.Fatalf("expected one client certificate, got %#v", client.TLSConfig) + } +} + +func testGigaChatReusesTLSClientWithCertificateUntilProviderReload(t *testing.T) { + t.Parallel() + + certPEM1, keyPEM1 := generateGigaChatTestCertificate(t) + certPEM2, keyPEM2 := generateGigaChatTestCertificate(t) + certFile := writeGigaChatTestFile(t, "client.pem", certPEM1) + keyFile := writeGigaChatTestFile(t, "client.key", keyPEM1) + + provider, err := NewGigaChatProvider(&schemas.ProviderConfig{}, nil) + if err != nil { + t.Fatalf("NewGigaChatProvider returned error: %v", err) + } + + keyConfig := &schemas.GigaChatKeyConfig{ + CertFile: certFile, + KeyFile: keyFile, + } + client, err := provider.getGigaChatTLSClient(provider.client, gigaChatTLSClientCacheDefault, keyConfig) + if err != nil { + t.Fatalf("getGigaChatTLSClient returned error: %v", err) + } + + if err := os.WriteFile(certFile, certPEM2, 0o600); err != nil { + t.Fatalf("failed to rotate client certificate file: %v", err) + } + if err := os.WriteFile(keyFile, keyPEM2, 0o600); err != nil { + t.Fatalf("failed to rotate client key file: %v", err) + } + + reusedClient, err := provider.getGigaChatTLSClient(provider.client, gigaChatTLSClientCacheDefault, keyConfig) + if err != nil { + t.Fatalf("getGigaChatTLSClient returned error after certificate rotation: %v", err) + } + if reusedClient != client { + t.Fatal("expected cached TLS client to be reused after certificate rotation") + } + + if err := os.Remove(certFile); err != nil { + t.Fatalf("failed to remove client certificate file: %v", err) + } + if err := os.Remove(keyFile); err != nil { + t.Fatalf("failed to remove client key file: %v", err) + } + reusedClient, err = provider.getGigaChatTLSClient(provider.client, gigaChatTLSClientCacheDefault, keyConfig) + if err != nil { + t.Fatalf("cached TLS client lookup read removed certificate files: %v", err) + } + if reusedClient != client { + t.Fatal("expected cached TLS client to be reused after certificate file removal") + } + + reloadedProvider, err := NewGigaChatProvider(&schemas.ProviderConfig{}, nil) + if err != nil { + t.Fatalf("NewGigaChatProvider returned error: %v", err) + } + if _, err := reloadedProvider.getGigaChatTLSClient(reloadedProvider.client, gigaChatTLSClientCacheDefault, keyConfig); err == nil { + t.Fatal("expected provider reload to validate missing certificate files") + } else if !strings.Contains(err.Error(), "cert_file/key_file") { + t.Fatalf("unexpected missing certificate error: %v", err) + } +} + +func testGigaChatRejectsMissingCertificatePair(t *testing.T) { + t.Parallel() + + provider, err := NewGigaChatProvider(&schemas.ProviderConfig{}, nil) + if err != nil { + t.Fatalf("NewGigaChatProvider returned error: %v", err) + } + + _, err = buildGigaChatTLSClient(provider.client, &schemas.GigaChatKeyConfig{CertFile: "client.pem"}) + if err == nil { + t.Fatal("expected error, got nil") + } + if !strings.Contains(err.Error(), "cert_file and gigachat_key_config.key_file") { + t.Fatalf("unexpected error: %v", err) + } +} + +func generateGigaChatTestCertificate(t *testing.T) ([]byte, []byte) { + t.Helper() + + privateKey, err := rsa.GenerateKey(rand.Reader, 2048) + if err != nil { + t.Fatalf("failed to generate private key: %v", err) + } + template := &x509.Certificate{ + SerialNumber: big.NewInt(1), + Subject: pkix.Name{ + CommonName: "gigachat-test", + }, + NotBefore: time.Now().Add(-time.Hour), + NotAfter: time.Now().Add(time.Hour), + KeyUsage: x509.KeyUsageDigitalSignature | x509.KeyUsageKeyEncipherment | x509.KeyUsageCertSign, + ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageClientAuth}, + BasicConstraintsValid: true, + IsCA: true, + } + certDER, err := x509.CreateCertificate(rand.Reader, template, template, &privateKey.PublicKey, privateKey) + if err != nil { + t.Fatalf("failed to create certificate: %v", err) + } + keyDER := x509.MarshalPKCS1PrivateKey(privateKey) + + certPEM := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: certDER}) + keyPEM := pem.EncodeToMemory(&pem.Block{Type: "RSA PRIVATE KEY", Bytes: keyDER}) + return certPEM, keyPEM +} + +func writeGigaChatTestFile(t *testing.T, name string, contents []byte) string { + t.Helper() + + path := t.TempDir() + "/" + name + if err := os.WriteFile(path, contents, 0o600); err != nil { + t.Fatalf("failed to write test file: %v", err) + } + return path +} diff --git a/core/providers/gigachat/key_config_test.go b/core/providers/gigachat/key_config_test.go new file mode 100644 index 00000000000..e344b839ce7 --- /dev/null +++ b/core/providers/gigachat/key_config_test.go @@ -0,0 +1,209 @@ +package gigachat + +import ( + "encoding/json" + "strings" + "testing" + + "github.com/maximhq/bifrost/core/schemas" +) + +func TestGigaChatKeyConfigRedacted(t *testing.T) { + t.Parallel() + + config := &schemas.GigaChatKeyConfig{ + Credentials: schemas.NewSecretVar("secret-credentials"), + User: schemas.NewSecretVar("secret-user"), + Password: schemas.NewSecretVar("secret-password"), + AccessToken: schemas.NewSecretVar("secret-access-token"), + CertFile: "/secure/client.pem", + KeyFile: "/secure/client.key", + CABundleFile: "/secure/ca.pem", + BaseURL: "https://api.giga.chat", + AuthURL: "https://ngw.devices.sberbank.ru:9443/api/v2/oauth", + } + + redacted := config.Redacted() + data, err := json.Marshal(redacted) + if err != nil { + t.Fatalf("json.Marshal returned error: %v", err) + } + output := string(data) + + for _, secret := range []string{ + "secret-credentials", + "secret-user", + "secret-password", + "secret-access-token", + "/secure/client.pem", + "/secure/client.key", + "/secure/ca.pem", + } { + if strings.Contains(output, secret) { + t.Fatalf("redacted config leaked %q in %s", secret, output) + } + } + if !strings.Contains(output, "https://api.giga.chat") { + t.Fatalf("non-secret base_url should be preserved in %s", output) + } +} + +func TestGigaChatKeyConfigEnvVarAuthMaterial(t *testing.T) { + t.Parallel() + + config := &schemas.GigaChatKeyConfig{ + Credentials: schemas.NewSecretVar("env.MISSING_GIGACHAT_CREDENTIALS"), + } + config.CheckAndSetDefaults() + + if config.Scope != schemas.DefaultGigaChatScope { + t.Fatalf("scope mismatch: got %q, want %q", config.Scope, schemas.DefaultGigaChatScope) + } + if !config.HasAuthMaterial() { + t.Fatal("expected unresolved env var reference to count as configured auth material") + } +} + +func TestGigaChatKeyConfigAuthMaterial(t *testing.T) { + t.Parallel() + + testCases := []struct { + name string + config *schemas.GigaChatKeyConfig + wantAuth bool + wantTLS bool + wantMTLS bool + }{ + { + name: "NilConfig", + config: nil, + wantAuth: false, + wantTLS: false, + wantMTLS: false, + }, + { + name: "Credentials", + config: &schemas.GigaChatKeyConfig{ + Credentials: schemas.NewSecretVar("env.GIGACHAT_CREDENTIALS"), + }, + wantAuth: true, + wantTLS: false, + wantMTLS: false, + }, + { + name: "AccessToken", + config: &schemas.GigaChatKeyConfig{ + AccessToken: schemas.NewSecretVar("env.GIGACHAT_ACCESS_TOKEN"), + }, + wantAuth: true, + wantTLS: false, + wantMTLS: false, + }, + { + name: "UserPassword", + config: &schemas.GigaChatKeyConfig{ + User: schemas.NewSecretVar("env.GIGACHAT_USER"), + Password: schemas.NewSecretVar("env.GIGACHAT_PASSWORD"), + }, + wantAuth: true, + wantTLS: false, + wantMTLS: false, + }, + { + name: "ClientCertificatePair", + config: &schemas.GigaChatKeyConfig{ + CertFile: "/secure/client.pem", + KeyFile: "/secure/client.key", + }, + wantAuth: false, + wantTLS: true, + wantMTLS: true, + }, + { + name: "CABundle", + config: &schemas.GigaChatKeyConfig{ + CABundleFile: "/secure/ca.pem", + }, + wantAuth: false, + wantTLS: true, + wantMTLS: false, + }, + { + name: "CredentialsWithTLS", + config: &schemas.GigaChatKeyConfig{ + Credentials: schemas.NewSecretVar("env.GIGACHAT_CREDENTIALS"), + CertFile: "/secure/client.pem", + KeyFile: "/secure/client.key", + CABundleFile: "/secure/ca.pem", + }, + wantAuth: true, + wantTLS: true, + wantMTLS: true, + }, + } + + for _, testCase := range testCases { + testCase := testCase + t.Run(testCase.name, func(t *testing.T) { + t.Parallel() + + if got := testCase.config.HasAuthMaterial(); got != testCase.wantAuth { + t.Fatalf("HasAuthMaterial mismatch: got %v, want %v", got, testCase.wantAuth) + } + if got := testCase.config.HasTLSMaterial(); got != testCase.wantTLS { + t.Fatalf("HasTLSMaterial mismatch: got %v, want %v", got, testCase.wantTLS) + } + if got := testCase.config.HasClientCertificateMaterial(); got != testCase.wantMTLS { + t.Fatalf("HasClientCertificateMaterial mismatch: got %v, want %v", got, testCase.wantMTLS) + } + }) + } +} + +func TestGigaChatKeyConfigValidate(t *testing.T) { + t.Parallel() + + t.Run("DefaultsScope", func(t *testing.T) { + t.Parallel() + + config := &schemas.GigaChatKeyConfig{ + Credentials: schemas.NewSecretVar("env.GIGACHAT_CREDENTIALS"), + } + if err := config.Validate(); err != nil { + t.Fatalf("Validate returned error: %v", err) + } + if config.Scope != schemas.DefaultGigaChatScope { + t.Fatalf("scope mismatch: got %q, want %q", config.Scope, schemas.DefaultGigaChatScope) + } + }) + + t.Run("PartialUserPasswordRejected", func(t *testing.T) { + t.Parallel() + + config := &schemas.GigaChatKeyConfig{ + User: schemas.NewSecretVar("env.GIGACHAT_USER"), + } + err := config.Validate() + if err == nil { + t.Fatal("expected validation error") + } + if !strings.Contains(err.Error(), "user and gigachat_key_config.password") { + t.Fatalf("unexpected error: %v", err) + } + }) + + t.Run("PartialCertificatePairRejected", func(t *testing.T) { + t.Parallel() + + config := &schemas.GigaChatKeyConfig{ + CertFile: "/secure/client.pem", + } + err := config.Validate() + if err == nil { + t.Fatal("expected validation error") + } + if !strings.Contains(err.Error(), "cert_file and gigachat_key_config.key_file") { + t.Fatalf("unexpected error: %v", err) + } + }) +} diff --git a/core/providers/gigachat/models_test.go b/core/providers/gigachat/models_test.go new file mode 100644 index 00000000000..f326736839b --- /dev/null +++ b/core/providers/gigachat/models_test.go @@ -0,0 +1,296 @@ +package gigachat + +import ( + "fmt" + "net/http" + "net/http/httptest" + "strings" + "sync/atomic" + "testing" + + "github.com/maximhq/bifrost/core/schemas" +) + +func testGigaChatListModels(t *testing.T) { + t.Parallel() + + t.Run("ConverterMapsResponse", testGigaChatListModelsConverterMapsResponse) + t.Run("ConverterFiltersAndAliases", testGigaChatListModelsConverterFiltersAndAliases) + t.Run("ExecutesWithOAuthToken", testGigaChatListModelsExecutesWithOAuthToken) + t.Run("MapsProviderErrors", testGigaChatListModelsMapsProviderErrors) + t.Run("RefreshesTokenAfterUnauthorized", testGigaChatListModelsRefreshesTokenAfterUnauthorized) +} + +func TestGigaChatSupportedMethods(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + modelType string + want []string + }{ + { + name: "chat includes responses", + modelType: "chat", + want: []string{ + string(schemas.ChatCompletionRequest), + string(schemas.ChatCompletionStreamRequest), + string(schemas.ResponsesRequest), + string(schemas.ResponsesStreamRequest), + }, + }, + { + name: "embedder only supports embeddings", + modelType: "embedder", + want: []string{string(schemas.EmbeddingRequest)}, + }, + { + name: "unknown has no advertised methods", + modelType: "reranker", + want: nil, + }, + } + + for _, tt := range tests { + tt := tt + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + if got := toGigaChatSupportedMethods(tt.modelType); fmt.Sprint(got) != fmt.Sprint(tt.want) { + t.Fatalf("toGigaChatSupportedMethods(%q) = %#v, want %#v", tt.modelType, got, tt.want) + } + }) + } +} + +func testGigaChatListModelsConverterMapsResponse(t *testing.T) { + t.Parallel() + + response := &GigaChatListModelsResponse{ + Object: "list", + Data: []GigaChatModel{ + {ID: "GigaChat", Object: "model", OwnedBy: "salutedevices", Type: "chat"}, + {ID: "Embeddings", Object: "model", OwnedBy: "salutedevices", Type: "embedder"}, + {ID: " ", Object: "model", OwnedBy: "salutedevices", Type: "chat"}, + }, + } + + converted := response.ToBifrostListModelsResponse(schemas.GigaChat, schemas.WhiteList{"*"}, nil, nil, false) + if converted == nil { + t.Fatal("expected response, got nil") + } + if len(converted.Data) != 2 { + t.Fatalf("model count mismatch: got %d, want 2", len(converted.Data)) + } + if converted.Data[0].ID != "gigachat/GigaChat" { + t.Fatalf("model id mismatch: got %q", converted.Data[0].ID) + } + if converted.Data[0].OwnedBy == nil || *converted.Data[0].OwnedBy != "salutedevices" { + t.Fatalf("owned_by mismatch: %#v", converted.Data[0].OwnedBy) + } + wantMethods := []string{ + string(schemas.ChatCompletionRequest), + string(schemas.ChatCompletionStreamRequest), + string(schemas.ResponsesRequest), + string(schemas.ResponsesStreamRequest), + } + if fmt.Sprint(converted.Data[0].SupportedMethods) != fmt.Sprint(wantMethods) { + t.Fatalf("supported methods mismatch: got %#v, want %#v", converted.Data[0].SupportedMethods, wantMethods) + } + wantEmbeddingMethods := []string{string(schemas.EmbeddingRequest)} + if fmt.Sprint(converted.Data[1].SupportedMethods) != fmt.Sprint(wantEmbeddingMethods) { + t.Fatalf("embedder supported methods mismatch: got %#v, want %#v", converted.Data[1].SupportedMethods, wantEmbeddingMethods) + } +} + +func testGigaChatListModelsConverterFiltersAndAliases(t *testing.T) { + t.Parallel() + + response := &GigaChatListModelsResponse{ + Data: []GigaChatModel{ + {ID: "GigaChat", Object: "model", OwnedBy: "salutedevices", Type: "chat"}, + {ID: "GigaChat-Pro", Object: "model", OwnedBy: "salutedevices", Type: "chat"}, + }, + } + + converted := response.ToBifrostListModelsResponse( + schemas.GigaChat, + schemas.WhiteList{"pro-alias"}, + nil, + schemas.KeyAliases{"pro-alias": {ModelID: "GigaChat-Pro"}}, + false, + ) + if converted == nil { + t.Fatal("expected response, got nil") + } + if len(converted.Data) != 1 { + t.Fatalf("model count mismatch: got %d, want 1", len(converted.Data)) + } + if converted.Data[0].ID != "gigachat/pro-alias" { + t.Fatalf("alias model id mismatch: got %q", converted.Data[0].ID) + } + if converted.Data[0].Alias == nil || *converted.Data[0].Alias != "GigaChat-Pro" { + t.Fatalf("alias mismatch: %#v", converted.Data[0].Alias) + } +} + +func testGigaChatListModelsExecutesWithOAuthToken(t *testing.T) { + t.Parallel() + + var tokenRequests atomic.Int32 + var modelRequests atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) { + switch request.URL.Path { + case "/oauth": + tokenRequests.Add(1) + if got := request.Header.Get("Authorization"); got != "Basic super-secret-credentials" { + t.Fatalf("token authorization header mismatch: got %q", got) + } + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"access_token":"models-access-token","expires_at":1893456000}`)) + case "/v1/models": + modelRequests.Add(1) + if request.Method != http.MethodGet { + t.Fatalf("method mismatch: got %s, want GET", request.Method) + } + if got := request.Header.Get("Authorization"); got != "Bearer models-access-token" { + t.Fatalf("models authorization header mismatch: got %q", got) + } + if strings.Contains(request.Header.Get("Authorization"), "super-secret-credentials") { + t.Fatal("models request leaked OAuth credentials") + } + if got := request.Header.Get("Accept"); got != "application/json" { + t.Fatalf("accept header mismatch: got %q", got) + } + if got := request.Header.Get(gigaChatUserAgentHeader); got != gigaChatUserAgent { + t.Fatalf("user-agent header mismatch: got %q", got) + } + w.Header().Set("Content-Type", "application/json") + w.Header().Set("X-Request-ID", "models-request-id") + _, _ = w.Write([]byte(`{"object":"list","data":[{"id":"GigaChat","object":"model","owned_by":"salutedevices","type":"chat"},{"id":"GigaChat-Pro","object":"model","owned_by":"salutedevices","type":"chat"}]}`)) + default: + t.Fatalf("unexpected path: %s", request.URL.Path) + } + })) + defer server.Close() + + provider := newTestGigaChatChatProvider(t, server.URL) + provider.sendBackRawResponse = true + key := testGigaChatOAuthKey(server.URL+"/oauth", "", "super-secret-credentials") + key.ID = "gigachat-key" + key.Models = schemas.WhiteList{"*"} + + ctx := testBifrostContext() + response, bifrostErr := provider.ListModels(ctx, []schemas.Key{key}, &schemas.BifrostListModelsRequest{Provider: schemas.GigaChat}) + if bifrostErr != nil { + t.Fatalf("ListModels returned error: %v", bifrostErr) + } + if tokenRequests.Load() != 1 { + t.Fatalf("token request count mismatch: got %d, want 1", tokenRequests.Load()) + } + if modelRequests.Load() != 1 { + t.Fatalf("model request count mismatch: got %d, want 1", modelRequests.Load()) + } + if len(response.Data) != 2 { + t.Fatalf("model count mismatch: got %d, want 2", len(response.Data)) + } + if response.Data[0].ID != "gigachat/GigaChat" || response.Data[1].ID != "gigachat/GigaChat-Pro" { + t.Fatalf("models mismatch: %#v", response.Data) + } + if response.ExtraFields.RawResponse == nil { + t.Fatal("expected raw response to be preserved") + } + if len(response.KeyStatuses) != 1 || response.KeyStatuses[0].Status != schemas.KeyStatusSuccess { + t.Fatalf("key status mismatch: %#v", response.KeyStatuses) + } + if got := ctx.Value(schemas.BifrostContextKeyProviderResponseHeaders); got == nil { + t.Fatal("provider response headers were not stored in context") + } +} + +func testGigaChatListModelsMapsProviderErrors(t *testing.T) { + t.Parallel() + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusBadRequest) + _, _ = w.Write([]byte(`{"status":400,"code":123,"message":"bad models request"}`)) + })) + defer server.Close() + + provider := newTestGigaChatChatProvider(t, server.URL) + key := testGigaChatAccessTokenKey("provider-error-token") + key.ID = "gigachat-key" + key.Models = schemas.WhiteList{"*"} + + response, bifrostErr := provider.ListModels(testBifrostContext(), []schemas.Key{key}, &schemas.BifrostListModelsRequest{Provider: schemas.GigaChat}) + if response != nil { + t.Fatalf("expected nil response, got %#v", response) + } + if bifrostErr == nil { + t.Fatal("expected provider error, got nil") + } + if bifrostErr.StatusCode == nil || *bifrostErr.StatusCode != http.StatusBadRequest { + t.Fatalf("status mismatch: %#v", bifrostErr.StatusCode) + } + if bifrostErr.Error == nil || bifrostErr.Error.Message != "bad models request" { + t.Fatalf("message mismatch: %#v", bifrostErr.Error) + } + if bifrostErr.Error.Code == nil || *bifrostErr.Error.Code != "123" { + t.Fatalf("code mismatch: %#v", bifrostErr.Error) + } + if len(bifrostErr.ExtraFields.KeyStatuses) != 1 || bifrostErr.ExtraFields.KeyStatuses[0].Status != schemas.KeyStatusListModelsFailed { + t.Fatalf("key status mismatch: %#v", bifrostErr.ExtraFields.KeyStatuses) + } + assertNoGigaChatSecretLeak(t, bifrostErr.String()) +} + +func testGigaChatListModelsRefreshesTokenAfterUnauthorized(t *testing.T) { + t.Parallel() + + var tokenRequests atomic.Int32 + var modelRequests atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) { + switch request.URL.Path { + case "/oauth": + tokenIndex := tokenRequests.Add(1) + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(fmt.Sprintf(`{"access_token":"models-token-%d","expires_at":1893456000}`, tokenIndex))) + case "/v1/models": + modelIndex := modelRequests.Add(1) + wantAuthorization := fmt.Sprintf("Bearer models-token-%d", modelIndex) + if got := request.Header.Get("Authorization"); got != wantAuthorization { + t.Fatalf("authorization header mismatch on request %d: got %q, want %q", modelIndex, got, wantAuthorization) + } + w.Header().Set("Content-Type", "application/json") + if modelIndex == 1 { + w.WriteHeader(http.StatusUnauthorized) + _, _ = w.Write([]byte(`{"status":401,"message":"expired token"}`)) + return + } + _, _ = w.Write([]byte(`{"object":"list","data":[{"id":"GigaChat","object":"model","owned_by":"salutedevices","type":"chat"}]}`)) + default: + t.Fatalf("unexpected path: %s", request.URL.Path) + } + })) + defer server.Close() + + provider := newTestGigaChatChatProvider(t, server.URL) + key := testGigaChatOAuthKey(server.URL+"/oauth", "", "test-credentials") + key.ID = "gigachat-key" + key.Models = schemas.WhiteList{"*"} + + response, bifrostErr := provider.ListModels(testBifrostContext(), []schemas.Key{key}, &schemas.BifrostListModelsRequest{Provider: schemas.GigaChat}) + if bifrostErr != nil { + t.Fatalf("ListModels returned error: %v", bifrostErr) + } + if response == nil || len(response.Data) != 1 || response.Data[0].ID != "gigachat/GigaChat" { + t.Fatalf("unexpected response: %#v", response) + } + if tokenRequests.Load() != 2 { + t.Fatalf("token request count mismatch: got %d, want 2", tokenRequests.Load()) + } + if modelRequests.Load() != 2 { + t.Fatalf("model request count mismatch: got %d, want 2", modelRequests.Load()) + } +} diff --git a/core/providers/gigachat/responses_test.go b/core/providers/gigachat/responses_test.go new file mode 100644 index 00000000000..1c3143da5d8 --- /dev/null +++ b/core/providers/gigachat/responses_test.go @@ -0,0 +1,2720 @@ +package gigachat + +import ( + "context" + "encoding/base64" + "encoding/json" + "fmt" + "io" + "net/http" + "net/http/httptest" + "strings" + "sync/atomic" + "testing" + "time" + + providerUtils "github.com/maximhq/bifrost/core/providers/utils" + "github.com/maximhq/bifrost/core/schemas" +) + +func TestGigaChatResponsesRequestConversion(t *testing.T) { + testGigaChatResponsesRequestConversion(t) +} + +func TestGigaChatResponses(t *testing.T) { + testGigaChatResponses(t) +} + +func TestGigaChatResponsesStream(t *testing.T) { + testGigaChatResponsesStream(t) +} + +func TestGigaChatResponsesAttachmentHelpers(t *testing.T) { + t.Run("RemoteImageURL", testGigaChatResponsesRemoteImageURLUploadHelper) + t.Run("RemoteFileURL", testGigaChatResponsesRemoteFileURLUploadHelper) +} + +func testGigaChatResponsesRequestConversion(t *testing.T) { + t.Parallel() + + t.Run("SimpleTextInput", testGigaChatResponsesSimpleTextInput) + t.Run("InstructionsAndMultiTurnInput", testGigaChatResponsesInstructionsAndMultiTurnInput) + t.Run("FileInputReference", testGigaChatResponsesFileInputReference) + t.Run("ImageInputReference", testGigaChatResponsesImageInputReference) + t.Run("FunctionToolAndToolHistory", testGigaChatResponsesFunctionToolAndToolHistory) + t.Run("StructuredOutput", testGigaChatResponsesStructuredOutput) + t.Run("RejectsUnsupportedHostedTools", testGigaChatResponsesRejectsUnsupportedHostedTools) + t.Run("RejectsUnsupportedParams", testGigaChatResponsesRejectsUnsupportedParams) + t.Run("RejectsUnsupportedFileInputs", testGigaChatResponsesRejectsUnsupportedFileInputs) + t.Run("FunctionCallOutputUsesCallIDAsToolsStateID", testGigaChatResponsesFunctionCallOutputUsesCallIDAsToolsStateID) + t.Run("FunctionCallOutputInfersNameFromGeneratedCallID", testGigaChatResponsesFunctionCallOutputInfersNameFromGeneratedCallID) + t.Run("ThreadStorage", testGigaChatResponsesThreadStorage) +} + +func testGigaChatResponses(t *testing.T) { + t.Parallel() + + t.Run("ConverterMapsTextAndUsage", testGigaChatResponsesConverterMapsTextAndUsage) + t.Run("ConverterMapsImageFileOutput", testGigaChatResponsesConverterMapsImageFileOutput) + t.Run("ConverterMapsWebSearchSources", testGigaChatResponsesConverterMapsWebSearchSources) + t.Run("ConverterMapsReasoningRole", testGigaChatResponsesConverterMapsReasoningRole) + t.Run("ConverterMapsToolCall", testGigaChatResponsesConverterMapsToolCall) + t.Run("ConverterMapsFunctionResultWithOriginatingCallID", testGigaChatResponsesConverterMapsFunctionResultWithOriginatingCallID) + t.Run("ConverterUsesUniqueCallIDsUnderSharedToolsStateID", testGigaChatResponsesConverterUsesUniqueCallIDsUnderSharedToolsStateID) + t.Run("ConverterUsesToolStateIDAliasAsCallID", testGigaChatResponsesConverterUsesToolStateIDAliasAsCallID) + t.Run("ConverterFallsBackToResponseToolsStateID", testGigaChatResponsesConverterFallsBackToResponseToolsStateID) + t.Run("ConverterPreservesOrdinaryMessageToolStateID", testGigaChatResponsesConverterPreservesOrdinaryMessageToolStateID) + t.Run("ConverterMapsThreadStorage", testGigaChatResponsesConverterMapsThreadStorage) + t.Run("ExecutesWithOAuthToken", testGigaChatResponsesExecutesWithOAuthToken) + t.Run("UploadsInputImageAttachment", testGigaChatResponsesUploadsInputImageAttachment) + t.Run("UploadsInlineFileAttachment", testGigaChatResponsesUploadsInlineFileAttachment) + t.Run("ReusesUploadedAttachmentAfterBackendError", testGigaChatResponsesReusesUploadedAttachmentAfterBackendError) + t.Run("ReusesCompletedUploadsAfterPartialAttachmentFailure", testGigaChatResponsesReusesCompletedUploadsAfterPartialAttachmentFailure) + t.Run("MapsProviderErrors", testGigaChatResponsesMapsProviderErrors) + t.Run("RefreshesTokenAfterUnauthorized", testGigaChatResponsesRefreshesTokenAfterUnauthorized) +} + +func testGigaChatResponsesStream(t *testing.T) { + t.Parallel() + + t.Run("TextDeltasAndUsage", testGigaChatResponsesStreamTextDeltasAndUsage) + t.Run("ReasoningDeltas", testGigaChatResponsesStreamReasoningDeltas) + t.Run("ToolCallDeltas", testGigaChatResponsesStreamToolCallDeltas) + t.Run("ClosesOnMessageDoneEvent", testGigaChatResponsesStreamClosesOnMessageDoneEvent) + t.Run("MapsErrorEvents", testGigaChatResponsesStreamMapsErrorEvents) + t.Run("HandlesContextCancellation", testGigaChatResponsesStreamHandlesContextCancellation) + t.Run("PassthroughResponseOwnedByLargeReader", testGigaChatResponsesStreamPassthroughResponseOwnedByLargeReader) +} + +func testGigaChatResponsesSimpleTextInput(t *testing.T) { + t.Parallel() + + temperature := 0.2 + topP := 0.8 + maxOutputTokens := 256 + topLogProbs := 3 + request := &schemas.BifrostResponsesRequest{ + Model: "GigaChat-2", + Input: []schemas.ResponsesMessage{{ + Role: schemas.Ptr(schemas.ResponsesInputMessageRoleUser), + Content: &schemas.ResponsesMessageContent{ContentStr: schemas.Ptr("hello")}, + }}, + Params: &schemas.ResponsesParameters{ + Temperature: &temperature, + TopP: &topP, + MaxOutputTokens: &maxOutputTokens, + TopLogProbs: &topLogProbs, + }, + } + + gigaChatReq, err := ToGigaChatResponsesRequest(request) + if err != nil { + t.Fatalf("ToGigaChatResponsesRequest returned error: %v", err) + } + if gigaChatReq.Model != "GigaChat-2" { + t.Fatalf("model mismatch: got %q", gigaChatReq.Model) + } + if len(gigaChatReq.Messages) != 1 { + t.Fatalf("message count mismatch: got %d", len(gigaChatReq.Messages)) + } + if got := gigaChatReq.Messages[0].Role; got != "user" { + t.Fatalf("role mismatch: got %q", got) + } + if got := *gigaChatReq.Messages[0].Content[0].Text; got != "hello" { + t.Fatalf("content mismatch: got %q", got) + } + if gigaChatReq.ModelOptions == nil || + gigaChatReq.ModelOptions.Temperature == nil || *gigaChatReq.ModelOptions.Temperature != temperature || + gigaChatReq.ModelOptions.TopP == nil || *gigaChatReq.ModelOptions.TopP != topP || + gigaChatReq.ModelOptions.MaxTokens == nil || *gigaChatReq.ModelOptions.MaxTokens != maxOutputTokens || + gigaChatReq.ModelOptions.TopLogProbs == nil || *gigaChatReq.ModelOptions.TopLogProbs != topLogProbs { + t.Fatalf("model options mismatch: %#v", gigaChatReq.ModelOptions) + } + + body, err := json.Marshal(gigaChatReq) + if err != nil { + t.Fatalf("failed to marshal GigaChat request: %v", err) + } + if strings.Contains(string(body), `"stream"`) { + t.Fatalf("non-streaming v2 request should omit stream, got %s", body) + } +} + +func testGigaChatResponsesInstructionsAndMultiTurnInput(t *testing.T) { + t.Parallel() + + instructions := "Answer in Russian." + request := &schemas.BifrostResponsesRequest{ + Model: "GigaChat-2-Pro", + Input: []schemas.ResponsesMessage{ + { + Role: schemas.Ptr(schemas.ResponsesInputMessageRoleSystem), + Content: &schemas.ResponsesMessageContent{ContentStr: schemas.Ptr("Be concise.")}, + }, + { + Role: schemas.Ptr(schemas.ResponsesInputMessageRoleAssistant), + Content: &schemas.ResponsesMessageContent{ContentBlocks: []schemas.ResponsesMessageContentBlock{{ + Type: schemas.ResponsesOutputMessageContentTypeText, + Text: schemas.Ptr("Previous answer."), + }}}, + }, + { + Role: schemas.Ptr(schemas.ResponsesInputMessageRoleUser), + Content: &schemas.ResponsesMessageContent{ContentBlocks: []schemas.ResponsesMessageContentBlock{{ + Type: schemas.ResponsesInputMessageContentBlockTypeText, + Text: schemas.Ptr("Continue."), + }}}, + }, + }, + Params: &schemas.ResponsesParameters{Instructions: &instructions}, + } + + gigaChatReq, err := ToGigaChatResponsesRequest(request) + if err != nil { + t.Fatalf("ToGigaChatResponsesRequest returned error: %v", err) + } + if len(gigaChatReq.Messages) != 4 { + t.Fatalf("message count mismatch: got %d", len(gigaChatReq.Messages)) + } + wantRoles := []string{"system", "system", "assistant", "user"} + wantText := []string{"Answer in Russian.", "Be concise.", "Previous answer.", "Continue."} + for index := range wantRoles { + if got := gigaChatReq.Messages[index].Role; got != wantRoles[index] { + t.Fatalf("message %d role mismatch: got %q, want %q", index, got, wantRoles[index]) + } + if got := *gigaChatReq.Messages[index].Content[0].Text; got != wantText[index] { + t.Fatalf("message %d text mismatch: got %q, want %q", index, got, wantText[index]) + } + } +} + +func testGigaChatResponsesFileInputReference(t *testing.T) { + t.Parallel() + + fileID := " file-document " + mime := " application/pdf " + filename := "document.pdf" + request := &schemas.BifrostResponsesRequest{ + Model: "GigaChat-2-Pro", + Input: []schemas.ResponsesMessage{{ + Role: schemas.Ptr(schemas.ResponsesInputMessageRoleUser), + Content: &schemas.ResponsesMessageContent{ContentBlocks: []schemas.ResponsesMessageContentBlock{ + { + Type: schemas.ResponsesInputMessageContentBlockTypeText, + Text: schemas.Ptr("Summarize this document."), + }, + { + Type: schemas.ResponsesInputMessageContentBlockTypeFile, + FileID: &fileID, + ResponsesInputMessageContentBlockFile: &schemas.ResponsesInputMessageContentBlockFile{ + Filename: &filename, + FileType: &mime, + }, + }, + }}, + }}, + } + + gigaChatReq, err := ToGigaChatResponsesRequest(request) + if err != nil { + t.Fatalf("ToGigaChatResponsesRequest returned error: %v", err) + } + if len(gigaChatReq.Messages) != 1 || len(gigaChatReq.Messages[0].Content) != 2 { + t.Fatalf("content parts mismatch: %#v", gigaChatReq.Messages) + } + files := gigaChatReq.Messages[0].Content[1].Files + if len(files) != 1 { + t.Fatalf("file refs mismatch: %#v", files) + } + if files[0].ID != "file-document" { + t.Fatalf("file id mismatch: got %q", files[0].ID) + } + if files[0].MIME == nil || *files[0].MIME != "application/pdf" { + t.Fatalf("file mime mismatch: %#v", files[0].MIME) + } + if files[0].Target != nil { + t.Fatalf("target should be omitted without a Bifrost source field, got %#v", files[0].Target) + } + + body, err := json.Marshal(gigaChatReq) + if err != nil { + t.Fatalf("failed to marshal GigaChat request: %v", err) + } + if !strings.Contains(string(body), `"files":[{"id":"file-document","mime":"application/pdf"}]`) { + t.Fatalf("request body should include GigaChat file reference, got %s", body) + } + if strings.Contains(string(body), filename) { + t.Fatalf("filename has no GigaChat v2 file content target mapping and should be omitted, got %s", body) + } +} + +func testGigaChatResponsesImageInputReference(t *testing.T) { + t.Parallel() + + fileID := " image-file " + request := &schemas.BifrostResponsesRequest{ + Model: "GigaChat-2-Pro", + Input: []schemas.ResponsesMessage{{ + Role: schemas.Ptr(schemas.ResponsesInputMessageRoleUser), + Content: &schemas.ResponsesMessageContent{ContentBlocks: []schemas.ResponsesMessageContentBlock{ + { + Type: schemas.ResponsesInputMessageContentBlockTypeText, + Text: schemas.Ptr("Describe this image."), + }, + { + Type: schemas.ResponsesInputMessageContentBlockTypeImage, + FileID: &fileID, + }, + }}, + }}, + } + + gigaChatReq, err := ToGigaChatResponsesRequest(request) + if err != nil { + t.Fatalf("ToGigaChatResponsesRequest returned error: %v", err) + } + if len(gigaChatReq.Messages) != 1 || len(gigaChatReq.Messages[0].Content) != 2 { + t.Fatalf("content parts mismatch: %#v", gigaChatReq.Messages) + } + files := gigaChatReq.Messages[0].Content[1].Files + if len(files) != 1 { + t.Fatalf("image file refs mismatch: %#v", files) + } + if files[0].ID != "image-file" { + t.Fatalf("image file id mismatch: got %q", files[0].ID) + } + if files[0].MIME != nil || files[0].Target != nil { + t.Fatalf("input_image file_id should omit mime and target, got %#v", files[0]) + } +} + +func testGigaChatResponsesFunctionToolAndToolHistory(t *testing.T) { + t.Parallel() + + toolName := "get_weather" + callID := "tools-state-weather" + arguments := `{"city":"Moscow"}` + toolOutput := `{"temperature":5}` + request := &schemas.BifrostResponsesRequest{ + Model: "GigaChat-2-Max", + Input: []schemas.ResponsesMessage{ + { + Role: schemas.Ptr(schemas.ResponsesInputMessageRoleUser), + Content: &schemas.ResponsesMessageContent{ContentStr: schemas.Ptr("Weather?")}, + }, + { + Type: schemas.Ptr(schemas.ResponsesMessageTypeFunctionCall), + ResponsesToolMessage: &schemas.ResponsesToolMessage{ + Name: &toolName, + CallID: &callID, + Arguments: &arguments, + }, + }, + { + Type: schemas.Ptr(schemas.ResponsesMessageTypeFunctionCallOutput), + ResponsesToolMessage: &schemas.ResponsesToolMessage{ + Name: &toolName, + CallID: &callID, + Output: &schemas.ResponsesToolMessageOutputStruct{ + ResponsesToolCallOutputStr: &toolOutput, + }, + }, + }, + }, + Params: &schemas.ResponsesParameters{ + Tools: []schemas.ResponsesTool{{ + Type: schemas.ResponsesToolTypeFunction, + Name: &toolName, + Description: schemas.Ptr("Gets current weather."), + ResponsesToolFunction: &schemas.ResponsesToolFunction{ + Parameters: mustGigaChatToolParameters(t, `{"type":"object","properties":{"city":{"type":"string"}},"required":["city"]}`), + }, + }}, + ToolChoice: &schemas.ResponsesToolChoice{ + ResponsesToolChoiceStruct: &schemas.ResponsesToolChoiceStruct{ + Type: schemas.ResponsesToolChoiceTypeFunction, + Name: &toolName, + }, + }, + }, + } + + gigaChatReq, err := ToGigaChatResponsesRequest(request) + if err != nil { + t.Fatalf("ToGigaChatResponsesRequest returned error: %v", err) + } + if len(gigaChatReq.Tools) != 1 || gigaChatReq.Tools[0].Functions == nil || len(gigaChatReq.Tools[0].Functions.Specifications) != 1 { + t.Fatalf("function tools mismatch: %#v", gigaChatReq.Tools) + } + specification := gigaChatReq.Tools[0].Functions.Specifications[0] + if specification.Name != toolName || specification.Description == nil || *specification.Description != "Gets current weather." { + t.Fatalf("function specification mismatch: %#v", specification) + } + if gigaChatReq.ToolConfig == nil || gigaChatReq.ToolConfig.Mode != "forced" || gigaChatReq.ToolConfig.FunctionName == nil || *gigaChatReq.ToolConfig.FunctionName != toolName { + t.Fatalf("tool config mismatch: %#v", gigaChatReq.ToolConfig) + } + if gigaChatReq.Messages[1].FunctionCall == nil { + t.Fatalf("expected function call message, got %#v", gigaChatReq.Messages[1]) + } + if gigaChatReq.Messages[1].ToolsStateID == nil || *gigaChatReq.Messages[1].ToolsStateID != callID { + t.Fatalf("function call tools_state_id mismatch: %#v", gigaChatReq.Messages[1].ToolsStateID) + } + argumentsMap, ok := gigaChatReq.Messages[1].FunctionCall.Arguments.(map[string]interface{}) + if !ok || argumentsMap["city"] != "Moscow" { + t.Fatalf("function arguments mismatch: %#v", gigaChatReq.Messages[1].FunctionCall.Arguments) + } + if gigaChatReq.Messages[2].ToolsStateID == nil || *gigaChatReq.Messages[2].ToolsStateID != callID { + t.Fatalf("function result tools_state_id mismatch: %#v", gigaChatReq.Messages[2].ToolsStateID) + } + if gigaChatReq.Messages[2].Content[0].FunctionResult == nil || gigaChatReq.Messages[2].Content[0].FunctionResult.Result != toolOutput { + t.Fatalf("function result mismatch: %#v", gigaChatReq.Messages[2].Content) + } +} + +func testGigaChatResponsesFunctionCallOutputUsesCallIDAsToolsStateID(t *testing.T) { + t.Parallel() + + toolName := "get_weather" + callID := "019e8282-bb13-73fc-bbe8-5f52856d166b__bifrost_fc_1" + toolOutput := `{"temperature":5}` + request := &schemas.BifrostResponsesRequest{ + Model: "GigaChat-2-Max", + Input: []schemas.ResponsesMessage{{ + Type: schemas.Ptr(schemas.ResponsesMessageTypeFunctionCallOutput), + ResponsesToolMessage: &schemas.ResponsesToolMessage{ + Name: &toolName, + CallID: &callID, + Output: &schemas.ResponsesToolMessageOutputStruct{ + ResponsesToolCallOutputStr: &toolOutput, + }, + }, + }}, + } + + gigaChatReq, err := ToGigaChatResponsesRequest(request) + if err != nil { + t.Fatalf("ToGigaChatResponsesRequest returned error: %v", err) + } + if len(gigaChatReq.Messages) != 1 { + t.Fatalf("message count mismatch: got %d", len(gigaChatReq.Messages)) + } + message := gigaChatReq.Messages[0] + if message.ToolsStateID == nil || *message.ToolsStateID != callID { + t.Fatalf("function_call_output tools_state_id mismatch: %#v", message.ToolsStateID) + } + if message.Content[0].FunctionResult == nil || message.Content[0].FunctionResult.Result != toolOutput { + t.Fatalf("function result mismatch: %#v", message.Content) + } +} + +func TestGigaChatResponsesFunctionCallOutputInfersNameFromPriorCallID(t *testing.T) { + t.Parallel() + + toolName := "get_weather" + callID := "019e8282-bb13-73fc-bbe8-5f52856d166b" + arguments := `{"city":"Moscow"}` + toolOutput := `{"temperature":5}` + request := &schemas.BifrostResponsesRequest{ + Model: "GigaChat-2-Max", + Input: []schemas.ResponsesMessage{ + { + Type: schemas.Ptr(schemas.ResponsesMessageTypeFunctionCall), + ResponsesToolMessage: &schemas.ResponsesToolMessage{ + Name: &toolName, + CallID: &callID, + Arguments: &arguments, + }, + }, + { + Type: schemas.Ptr(schemas.ResponsesMessageTypeFunctionCallOutput), + ResponsesToolMessage: &schemas.ResponsesToolMessage{ + CallID: &callID, + Output: &schemas.ResponsesToolMessageOutputStruct{ + ResponsesToolCallOutputStr: &toolOutput, + }, + }, + }, + }, + } + + gigaChatReq, err := ToGigaChatResponsesRequest(request) + if err != nil { + t.Fatalf("ToGigaChatResponsesRequest returned error: %v", err) + } + if len(gigaChatReq.Messages) != 2 { + t.Fatalf("message count mismatch: got %d", len(gigaChatReq.Messages)) + } + functionResult := gigaChatReq.Messages[1].Content[0].FunctionResult + if functionResult == nil { + t.Fatalf("function result missing: %#v", gigaChatReq.Messages[1].Content) + } + if functionResult.Name != toolName || functionResult.Result != toolOutput { + t.Fatalf("function result mismatch: %#v", functionResult) + } +} + +func testGigaChatResponsesFunctionCallOutputInfersNameFromGeneratedCallID(t *testing.T) { + t.Parallel() + + response := &GigaChatResponsesResponse{ + Model: "GigaChat-2-Max", + Messages: []GigaChatResponsesMessage{{ + Role: "assistant", + MessageID: schemas.Ptr("call-message"), + ToolsStateID: schemas.Ptr("tools-state-call"), + Content: []GigaChatResponsesContentPart{{ + FunctionCall: &GigaChatResponsesFunctionCall{ + Name: "get_weather", + Arguments: map[string]interface{}{"city": "Moscow"}, + }, + }}, + }}, + } + converted := ToBifrostResponsesResponse(schemas.GigaChat, response) + if converted == nil || len(converted.Output) != 1 || converted.Output[0].ResponsesToolMessage == nil { + t.Fatalf("converted output mismatch: %#v", converted) + } + + threadID := "thread-123" + toolOutput := `{"temperature":5}` + request := &schemas.BifrostResponsesRequest{ + Model: "GigaChat-2-Max", + Input: []schemas.ResponsesMessage{{ + Type: schemas.Ptr(schemas.ResponsesMessageTypeFunctionCallOutput), + ResponsesToolMessage: &schemas.ResponsesToolMessage{ + CallID: converted.Output[0].ResponsesToolMessage.CallID, + Output: &schemas.ResponsesToolMessageOutputStruct{ + ResponsesToolCallOutputStr: &toolOutput, + }, + }, + }}, + Params: &schemas.ResponsesParameters{ + PreviousResponseID: &threadID, + }, + } + + gigaChatReq, err := ToGigaChatResponsesRequest(request) + if err != nil { + t.Fatalf("ToGigaChatResponsesRequest returned error: %v", err) + } + if len(gigaChatReq.Messages) != 1 { + t.Fatalf("message count mismatch: got %d", len(gigaChatReq.Messages)) + } + if gigaChatReq.Model != "" { + t.Fatalf("previous_response_id request should omit provider model, got %q", gigaChatReq.Model) + } + message := gigaChatReq.Messages[0] + assertGigaChatResponsesToolStateID(t, message, "tools-state-call") + functionResult := message.Content[0].FunctionResult + if functionResult == nil || functionResult.Name != "get_weather" || functionResult.Result != toolOutput { + t.Fatalf("function result mismatch: %#v", message.Content) + } +} + +func testGigaChatResponsesThreadStorage(t *testing.T) { + t.Parallel() + + request := testGigaChatResponsesRequest() + gigaChatReq, err := ToGigaChatResponsesRequest(request) + if err != nil { + t.Fatalf("ToGigaChatResponsesRequest returned error: %v", err) + } + storage, ok := gigaChatReq.Storage.(*GigaChatResponsesStorage) + if !ok { + t.Fatalf("storage should default to an object, got %#v", gigaChatReq.Storage) + } + if storage.ThreadID != nil || len(storage.Metadata) != 0 { + t.Fatalf("default storage should be empty, got %#v", storage) + } + body, err := json.Marshal(gigaChatReq) + if err != nil { + t.Fatalf("failed to marshal GigaChat request: %v", err) + } + if !strings.Contains(string(body), `"storage":{}`) { + t.Fatalf("default storage object missing from request: %s", body) + } + if !strings.Contains(string(body), `"model":"GigaChat-2"`) { + t.Fatalf("initial thread request should include model, got %s", body) + } + + threadID := "thread-123" + metadata := map[string]any{"tenant": "test"} + request.Params = &schemas.ResponsesParameters{ + PreviousResponseID: &threadID, + Metadata: &metadata, + } + gigaChatReq, err = ToGigaChatResponsesRequest(request) + if err != nil { + t.Fatalf("ToGigaChatResponsesRequest with previous_response_id returned error: %v", err) + } + storage, ok = gigaChatReq.Storage.(*GigaChatResponsesStorage) + if !ok || storage.ThreadID == nil || *storage.ThreadID != threadID { + t.Fatalf("thread storage mismatch: %#v", gigaChatReq.Storage) + } + if storage.Metadata["tenant"] != "test" { + t.Fatalf("storage metadata mismatch: %#v", storage.Metadata) + } + if gigaChatReq.Model != "" { + t.Fatalf("previous_response_id request should omit provider model, got %q", gigaChatReq.Model) + } + body, err = json.Marshal(gigaChatReq) + if err != nil { + t.Fatalf("failed to marshal GigaChat request with previous_response_id: %v", err) + } + if strings.Contains(string(body), `"model"`) { + t.Fatalf("previous_response_id request should omit model from provider body, got %s", body) + } + if !strings.Contains(string(body), `"thread_id":"thread-123"`) { + t.Fatalf("previous_response_id request should include thread_id, got %s", body) + } + + request.Params = &schemas.ResponsesParameters{ + Conversation: &threadID, + } + gigaChatReq, err = ToGigaChatResponsesRequest(request) + if err != nil { + t.Fatalf("ToGigaChatResponsesRequest with conversation returned error: %v", err) + } + if gigaChatReq.Model != "" { + t.Fatalf("conversation request should omit provider model, got %q", gigaChatReq.Model) + } + + store := false + request.Params = &schemas.ResponsesParameters{Store: &store} + gigaChatReq, err = ToGigaChatResponsesRequest(request) + if err != nil { + t.Fatalf("ToGigaChatResponsesRequest with store=false returned error: %v", err) + } + if disabled, ok := gigaChatReq.Storage.(bool); !ok || disabled { + t.Fatalf("store=false should map to storage=false, got %#v", gigaChatReq.Storage) + } + + otherThreadID := "thread-456" + request.Params = &schemas.ResponsesParameters{ + Conversation: &threadID, + PreviousResponseID: &otherThreadID, + } + _, err = ToGigaChatResponsesRequest(request) + if err == nil || !strings.Contains(err.Error(), "same thread_id") { + t.Fatalf("expected thread id conflict error, got %v", err) + } +} + +func testGigaChatResponsesRemoteImageURLUploadHelper(t *testing.T) { + t.Parallel() + + imageURL := "https://cdn.example.com/assets/cat.png" + fetch := func(_ context.Context, resourceURL string) (string, string, error) { + if resourceURL != imageURL { + t.Fatalf("resource URL mismatch: got %q, want %q", resourceURL, imageURL) + } + return "image/png", base64.StdEncoding.EncodeToString([]byte("remote-image")), nil + } + + upload, err := gigaChatResponsesImageURLUpload(context.Background(), 0, schemas.ResponsesMessageContentBlock{ + Type: schemas.ResponsesInputMessageContentBlockTypeImage, + ResponsesInputMessageContentBlockImage: &schemas.ResponsesInputMessageContentBlockImage{ + ImageURL: &imageURL, + }, + }, fetch) + if err != nil { + t.Fatalf("gigaChatResponsesImageURLUpload returned error: %v", err) + } + if string(upload.file) != "remote-image" { + t.Fatalf("upload bytes mismatch: %q", upload.file) + } + if upload.filename != "cat.png" { + t.Fatalf("filename mismatch: got %q", upload.filename) + } + if upload.contentType != "image/png" { + t.Fatalf("content type mismatch: got %q", upload.contentType) + } +} + +func testGigaChatResponsesRemoteFileURLUploadHelper(t *testing.T) { + t.Parallel() + + fileURL := "https://cdn.example.com/docs/report" + filename := "report.pdf" + fetch := func(_ context.Context, resourceURL string) (string, string, error) { + if resourceURL != fileURL { + t.Fatalf("resource URL mismatch: got %q, want %q", resourceURL, fileURL) + } + return "application/pdf; charset=binary", base64.StdEncoding.EncodeToString([]byte("%PDF remote")), nil + } + + upload, err := gigaChatResponsesFileURLUpload(context.Background(), 1, &schemas.ResponsesInputMessageContentBlockFile{ + FileURL: &fileURL, + Filename: &filename, + }, fetch) + if err != nil { + t.Fatalf("gigaChatResponsesFileURLUpload returned error: %v", err) + } + if string(upload.file) != "%PDF remote" { + t.Fatalf("upload bytes mismatch: %q", upload.file) + } + if upload.filename != "report.pdf" { + t.Fatalf("filename mismatch: got %q", upload.filename) + } + if upload.contentType != "application/pdf" { + t.Fatalf("content type mismatch: got %q", upload.contentType) + } +} + +func testGigaChatResponsesRejectsUnsupportedFileInputs(t *testing.T) { + t.Parallel() + + fileID := "file-document" + fileData := "SGVsbG8=" + fileURL := "https://example.test/document.pdf" + cases := []struct { + name string + block schemas.ResponsesMessageContentBlock + wantErrSub string + }{ + { + name: "MissingFileID", + block: schemas.ResponsesMessageContentBlock{ + Type: schemas.ResponsesInputMessageContentBlockTypeFile, + FileID: schemas.Ptr(" "), + ResponsesInputMessageContentBlockFile: &schemas.ResponsesInputMessageContentBlockFile{ + FileType: schemas.Ptr("text/plain"), + }, + }, + wantErrSub: "requires file_id", + }, + { + name: "InlineFileData", + block: schemas.ResponsesMessageContentBlock{ + Type: schemas.ResponsesInputMessageContentBlockTypeFile, + FileID: &fileID, + ResponsesInputMessageContentBlockFile: &schemas.ResponsesInputMessageContentBlockFile{ + FileData: &fileData, + }, + }, + wantErrSub: "pre-uploaded file_id references only", + }, + { + name: "InlineFileURL", + block: schemas.ResponsesMessageContentBlock{ + Type: schemas.ResponsesInputMessageContentBlockTypeFile, + ResponsesInputMessageContentBlockFile: &schemas.ResponsesInputMessageContentBlockFile{ + FileURL: &fileURL, + }, + }, + wantErrSub: "pre-uploaded file_id references only", + }, + } + + for _, tc := range cases { + tc := tc + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + + request := &schemas.BifrostResponsesRequest{ + Model: "GigaChat-2", + Input: []schemas.ResponsesMessage{{ + Role: schemas.Ptr(schemas.ResponsesInputMessageRoleUser), + Content: &schemas.ResponsesMessageContent{ContentBlocks: []schemas.ResponsesMessageContentBlock{ + tc.block, + }}, + }}, + } + + _, err := ToGigaChatResponsesRequest(request) + if err == nil || !strings.Contains(err.Error(), tc.wantErrSub) { + t.Fatalf("expected %q error, got %v", tc.wantErrSub, err) + } + }) + } +} + +func testGigaChatResponsesStructuredOutput(t *testing.T) { + t.Parallel() + + strict := true + formatName := "WeatherAnswer" + formatDescription := "Weather response." + sourceSchema := schemas.NewOrderedMapFromPairs( + schemas.KV("type", "object"), + schemas.KV("properties", schemas.NewOrderedMapFromPairs( + schemas.KV("answer", schemas.NewOrderedMapFromPairs(schemas.KV("type", "string"))), + )), + schemas.KV("required", []string{"answer"}), + ) + request := &schemas.BifrostResponsesRequest{ + Model: "GigaChat-2-Pro", + Input: []schemas.ResponsesMessage{{ + Role: schemas.Ptr(schemas.ResponsesInputMessageRoleUser), + Content: &schemas.ResponsesMessageContent{ContentStr: schemas.Ptr("Return JSON.")}, + }}, + Params: &schemas.ResponsesParameters{ + Text: &schemas.ResponsesTextConfig{ + Format: &schemas.ResponsesTextConfigFormat{ + Type: "json_schema", + Name: &formatName, + Description: &formatDescription, + Strict: &strict, + JSONSchema: &schemas.ResponsesTextConfigFormatJSONSchema{ + Schema: &schemas.JSONSchemaOrBool{SchemaMap: sourceSchema}, + }, + }, + }, + }, + } + + gigaChatReq, err := ToGigaChatResponsesRequest(request) + if err != nil { + t.Fatalf("ToGigaChatResponsesRequest returned error: %v", err) + } + responseFormat := gigaChatReq.ModelOptions.ResponseFormat + if responseFormat == nil || responseFormat.Type != "json_schema" || responseFormat.Strict == nil || !*responseFormat.Strict { + t.Fatalf("response format mismatch: %#v", responseFormat) + } + schemaMap, ok := responseFormat.Schema.(*schemas.OrderedMap) + if !ok { + t.Fatalf("response schema has unexpected type: %#v", responseFormat.Schema) + } + title, _ := schemaMap.Get("title") + description, _ := schemaMap.Get("description") + if title != formatName || description != formatDescription { + t.Fatalf("schema metadata mismatch: %#v", schemaMap) + } + schemaType, _ := schemaMap.Get("type") + if schemaType != "object" { + t.Fatalf("schema type mismatch: %#v", schemaMap) + } + if _, exists := sourceSchema.Get("title"); exists { + t.Fatalf("source schema was mutated with title metadata: %#v", sourceSchema) + } + if _, exists := sourceSchema.Get("description"); exists { + t.Fatalf("source schema was mutated with description metadata: %#v", sourceSchema) + } +} + +func testGigaChatResponsesRejectsUnsupportedHostedTools(t *testing.T) { + t.Parallel() + + request := testGigaChatResponsesRequest() + request.Params = &schemas.ResponsesParameters{ + Tools: []schemas.ResponsesTool{{ + Type: schemas.ResponsesToolTypeFileSearch, + ResponsesToolFileSearch: &schemas.ResponsesToolFileSearch{ + VectorStoreIDs: []string{"vs_123"}, + }, + }}, + } + + _, err := ToGigaChatResponsesRequest(request) + if err == nil { + t.Fatal("expected unsupported hosted tool error, got nil") + } + if !strings.Contains(err.Error(), "does not support tool type") || !strings.Contains(err.Error(), "file_search") { + t.Fatalf("unexpected error: %v", err) + } +} + +func testGigaChatResponsesRejectsUnsupportedParams(t *testing.T) { + t.Parallel() + + parallelToolCalls := true + request := testGigaChatResponsesRequest() + request.Params = &schemas.ResponsesParameters{ParallelToolCalls: ¶llelToolCalls} + _, err := ToGigaChatResponsesRequest(request) + if err == nil || !strings.Contains(err.Error(), "parallel_tool_calls") { + t.Fatalf("expected parallel_tool_calls error, got %v", err) + } + + request = testGigaChatResponsesRequest() + request.Params = &schemas.ResponsesParameters{ + Text: &schemas.ResponsesTextConfig{ + Format: &schemas.ResponsesTextConfigFormat{Type: "json_object"}, + }, + } + _, err = ToGigaChatResponsesRequest(request) + if err == nil || !strings.Contains(err.Error(), "json_object") { + t.Fatalf("expected json_object error, got %v", err) + } + + request = &schemas.BifrostResponsesRequest{ + Model: "GigaChat-2", + Input: []schemas.ResponsesMessage{{ + Role: schemas.Ptr(schemas.ResponsesInputMessageRoleUser), + Content: &schemas.ResponsesMessageContent{ContentBlocks: []schemas.ResponsesMessageContentBlock{{ + Type: schemas.ResponsesInputMessageContentBlockTypeImage, + ResponsesInputMessageContentBlockImage: &schemas.ResponsesInputMessageContentBlockImage{ + ImageURL: schemas.Ptr("https://example.test/image.png"), + }, + }}}, + }}, + } + _, err = ToGigaChatResponsesRequest(request) + if err == nil || !strings.Contains(err.Error(), "input_image") { + t.Fatalf("expected input_image error, got %v", err) + } +} + +func testGigaChatResponsesConverterMapsTextAndUsage(t *testing.T) { + t.Parallel() + + response := &GigaChatResponsesResponse{ + MessageID: schemas.Ptr("resp-test"), + CreatedAt: 1700000000, + Model: "GigaChat-2", + Messages: []GigaChatResponsesMessage{{ + Role: "assistant", + MessageID: schemas.Ptr("msg-test"), + Content: []GigaChatResponsesContentPart{{ + Text: schemas.Ptr("Здравствуйте"), + }}, + FinishReason: schemas.Ptr("stop"), + }}, + Usage: &GigaChatChatUsage{ + InputTokens: 7, + OutputTokens: 3, + TotalTokens: 10, + InputTokensDetails: &GigaChatTokenDetails{ + CachedTokens: 2, + }, + }, + } + + converted := ToBifrostResponsesResponse(schemas.GigaChat, response) + if converted == nil { + t.Fatal("expected response, got nil") + } + if converted.ID == nil || *converted.ID != "resp-test" { + t.Fatalf("id mismatch: %#v", converted.ID) + } + if converted.Object != "response" || converted.CreatedAt != 1700000000 || converted.Model != "GigaChat-2" { + t.Fatalf("metadata mismatch: %#v", converted) + } + if converted.Status == nil || *converted.Status != "completed" { + t.Fatalf("status mismatch: %#v", converted.Status) + } + if converted.StopReason == nil || *converted.StopReason != "stop" { + t.Fatalf("stop reason mismatch: %#v", converted.StopReason) + } + if converted.Usage == nil || converted.Usage.InputTokens != 7 || converted.Usage.OutputTokens != 3 || converted.Usage.TotalTokens != 10 { + t.Fatalf("usage mismatch: %#v", converted.Usage) + } + if converted.Usage.InputTokensDetails == nil || converted.Usage.InputTokensDetails.CachedReadTokens != 2 { + t.Fatalf("cached tokens mismatch: %#v", converted.Usage.InputTokensDetails) + } + if len(converted.Output) != 1 { + t.Fatalf("output count mismatch: got %d", len(converted.Output)) + } + output := converted.Output[0] + if output.Type == nil || *output.Type != schemas.ResponsesMessageTypeMessage { + t.Fatalf("output type mismatch: %#v", output.Type) + } + if output.Role == nil || *output.Role != schemas.ResponsesInputMessageRoleAssistant { + t.Fatalf("output role mismatch: %#v", output.Role) + } + if output.Content == nil || len(output.Content.ContentBlocks) != 1 { + t.Fatalf("content mismatch: %#v", output.Content) + } + block := output.Content.ContentBlocks[0] + if block.Type != schemas.ResponsesOutputMessageContentTypeText || block.Text == nil || *block.Text != "Здравствуйте" { + t.Fatalf("text block mismatch: %#v", block) + } + if converted.ExtraFields.Provider != schemas.GigaChat { + t.Fatalf("provider mismatch: got %q, want %q", converted.ExtraFields.Provider, schemas.GigaChat) + } +} + +func testGigaChatResponsesConverterMapsImageFileOutput(t *testing.T) { + t.Parallel() + + fileID := "629ea825-963c-4178-bea4-1415c1d15a6e" + response := &GigaChatResponsesResponse{ + Model: "GigaChat-3-Ultra", + Messages: []GigaChatResponsesMessage{{ + Role: "assistant", + MessageID: schemas.Ptr("msg-image"), + Content: []GigaChatResponsesContentPart{ + { + Files: []GigaChatResponsesContentFile{{ + ID: fileID, + MIME: schemas.Ptr("image/jpeg"), + Target: schemas.Ptr("image"), + }}, + }, + { + Text: schemas.Ptr("вот красивая картинка с коровой в космосе."), + }, + }, + FinishReason: schemas.Ptr("stop"), + }}, + } + + converted := ToBifrostResponsesResponse(schemas.GigaChat, response) + if converted == nil { + t.Fatal("expected response, got nil") + } + if len(converted.Output) != 2 { + t.Fatalf("output count mismatch: got %d", len(converted.Output)) + } + + message := converted.Output[0] + if message.Type == nil || *message.Type != schemas.ResponsesMessageTypeMessage { + t.Fatalf("assistant output type mismatch: %#v", message.Type) + } + if message.Content == nil || len(message.Content.ContentBlocks) != 1 { + t.Fatalf("assistant content mismatch: %#v", message.Content) + } + textBlock := message.Content.ContentBlocks[0] + if textBlock.Text == nil || *textBlock.Text != "вот красивая картинка с коровой в космосе." { + t.Fatalf("assistant text mismatch: %#v", textBlock) + } + + imageCall := converted.Output[1] + if imageCall.ID == nil || *imageCall.ID != "ig_msg-image_0" { + t.Fatalf("image output id mismatch: %#v", imageCall.ID) + } + if imageCall.Type == nil || *imageCall.Type != schemas.ResponsesMessageTypeImageGenerationCall { + t.Fatalf("image output type mismatch: %#v", imageCall.Type) + } + if imageCall.Status == nil || *imageCall.Status != "completed" { + t.Fatalf("image output status mismatch: %#v", imageCall.Status) + } + if imageCall.ResponsesToolMessage == nil || imageCall.ResponsesToolMessage.ResponsesImageGenerationCall == nil { + t.Fatalf("image generation call missing: %#v", imageCall.ResponsesToolMessage) + } + if imageCall.ResponsesToolMessage.ResponsesImageGenerationCall.Result != fileID { + t.Fatalf("image generation result mismatch: %#v", imageCall.ResponsesToolMessage.ResponsesImageGenerationCall) + } + + raw, err := json.Marshal(imageCall) + if err != nil { + t.Fatalf("failed to marshal image output: %v", err) + } + var marshaled map[string]interface{} + if err := json.Unmarshal(raw, &marshaled); err != nil { + t.Fatalf("failed to unmarshal image output JSON: %v", err) + } + if marshaled["type"] != string(schemas.ResponsesMessageTypeImageGenerationCall) || marshaled["result"] != fileID { + t.Fatalf("image output JSON mismatch: %s", string(raw)) + } +} + +func testGigaChatResponsesConverterMapsWebSearchSources(t *testing.T) { + t.Parallel() + + text := "На сегодня курс доллара США к рублю, установленный Центральным банком России, составляет 72,56 RUB за 1 USD. [sources=[2, 4]]" + response := &GigaChatResponsesResponse{ + ThreadID: schemas.Ptr("377b9094-29c4-4711-99a0-69145bb764ec"), + MessageID: schemas.Ptr("ba4c614c-9c04-4eea-970d-ee22f1f0f94e"), + Model: "GigaChat-3-Ultra:32.3.18.5", + Messages: []GigaChatResponsesMessage{{ + Role: "assistant", + ToolStateID: schemas.Ptr("019e8cbd-1f65-77aa-9d24-400001ba0ccf"), + Content: []GigaChatResponsesContentPart{{ + Text: schemas.Ptr(text), + InlineData: map[string]interface{}{ + "images": []interface{}{}, + "sources": map[string]interface{}{ + "1": map[string]interface{}{ + "url": "https://cbr.ru/currency_base/daily/", + "title": "Официальные курсы валют на заданную дату, устанавливаемые...", + }, + "2": map[string]interface{}{ + "url": "https://www.vbr.ru/banki/kurs-valut/cbrf/usd/", + "title": "Курс доллара США к рублю на сегодня и завтра — Официальный...", + }, + "3": map[string]interface{}{ + "url": "https://www.banki.ru/products/currency/cash/moskva/", + "title": "Курсы валют в Москве на сегодня, выгодный курс обмена...", + }, + "4": map[string]interface{}{ + "url": "https://www.profinance.ru/cbrf/usd", + "title": "Курс доллара к рублю сегодня: онлайн графики и все котировки", + }, + "5": map[string]interface{}{ + "url": "https://news.ru/vlast/centrobank-rossii-ponizil-kurs-dollara", + "title": "Центробанк России понизил курс доллара", + }, + }, + }, + }}, + FinishReason: schemas.Ptr("stop"), + }}, + } + + converted := ToBifrostResponsesResponse(schemas.GigaChat, response) + if converted == nil { + t.Fatal("expected response, got nil") + } + if len(converted.Output) != 2 { + t.Fatalf("output count mismatch: got %d", len(converted.Output)) + } + + message := converted.Output[0] + if message.Type == nil || *message.Type != schemas.ResponsesMessageTypeMessage { + t.Fatalf("assistant output type mismatch: %#v", message.Type) + } + if message.Content == nil || len(message.Content.ContentBlocks) != 1 { + t.Fatalf("assistant content mismatch: %#v", message.Content) + } + block := message.Content.ContentBlocks[0] + if block.ResponsesOutputMessageContentText == nil { + t.Fatalf("text metadata missing: %#v", block) + } + annotations := block.ResponsesOutputMessageContentText.Annotations + if len(annotations) != 2 { + t.Fatalf("annotations count mismatch: got %d, annotations=%#v", len(annotations), annotations) + } + marker := "[sources=[2, 4]]" + startIndex := strings.Index(text, marker) + endIndex := startIndex + len(marker) + if annotations[0].Type != "url_citation" || annotations[0].URL == nil || *annotations[0].URL != "https://www.vbr.ru/banki/kurs-valut/cbrf/usd/" { + t.Fatalf("first annotation mismatch: %#v", annotations[0]) + } + if annotations[0].Title == nil || *annotations[0].Title != "Курс доллара США к рублю на сегодня и завтра — Официальный..." { + t.Fatalf("first annotation title mismatch: %#v", annotations[0].Title) + } + if annotations[0].StartIndex == nil || *annotations[0].StartIndex != startIndex || annotations[0].EndIndex == nil || *annotations[0].EndIndex != endIndex { + t.Fatalf("first annotation range mismatch: %#v", annotations[0]) + } + if annotations[1].Type != "url_citation" || annotations[1].URL == nil || *annotations[1].URL != "https://www.profinance.ru/cbrf/usd" { + t.Fatalf("second annotation mismatch: %#v", annotations[1]) + } + + webSearch := converted.Output[1] + if webSearch.ID == nil || *webSearch.ID != "ws_ba4c614c-9c04-4eea-970d-ee22f1f0f94e_0" { + t.Fatalf("web search output id mismatch: %#v", webSearch.ID) + } + if webSearch.Type == nil || *webSearch.Type != schemas.ResponsesMessageTypeWebSearchCall { + t.Fatalf("web search output type mismatch: %#v", webSearch.Type) + } + if webSearch.ResponsesToolMessage == nil || webSearch.ResponsesToolMessage.Action == nil || webSearch.ResponsesToolMessage.Action.ResponsesWebSearchToolCallAction == nil { + t.Fatalf("web search action missing: %#v", webSearch.ResponsesToolMessage) + } + action := webSearch.ResponsesToolMessage.Action.ResponsesWebSearchToolCallAction + if action.Type != "search" || len(action.Sources) != 5 { + t.Fatalf("web search sources mismatch: %#v", action) + } + if action.Sources[0].URL != "https://cbr.ru/currency_base/daily/" || action.Sources[4].URL != "https://news.ru/vlast/centrobank-rossii-ponizil-kurs-dollara" { + t.Fatalf("web search source ordering mismatch: %#v", action.Sources) + } +} + +func testGigaChatResponsesConverterMapsReasoningRole(t *testing.T) { + t.Parallel() + + response := &GigaChatResponsesResponse{ + CreatedAt: 1780306293, + Model: "GigaChat-2-Reasoning:2.0.29.05", + Messages: []GigaChatResponsesMessage{ + { + Role: "reasoning", + Content: []GigaChatResponsesContentPart{{ + Text: schemas.Ptr("...reasoning text..."), + }}, + }, + { + Role: "assistant", + Content: []GigaChatResponsesContentPart{{ + Text: schemas.Ptr("**Столица Франции — Париж.**"), + }}, + }, + }, + FinishReason: schemas.Ptr("stop"), + } + + converted := ToBifrostResponsesResponse(schemas.GigaChat, response) + if converted == nil { + t.Fatal("expected response, got nil") + } + if len(converted.Output) != 2 { + t.Fatalf("output count mismatch: got %d", len(converted.Output)) + } + + reasoning := converted.Output[0] + if reasoning.Type == nil || *reasoning.Type != schemas.ResponsesMessageTypeReasoning { + t.Fatalf("reasoning output type mismatch: %#v", reasoning.Type) + } + if reasoning.Role == nil || *reasoning.Role != schemas.ResponsesInputMessageRoleAssistant { + t.Fatalf("reasoning output role mismatch: %#v", reasoning.Role) + } + if reasoning.Status == nil || *reasoning.Status != "completed" { + t.Fatalf("reasoning status mismatch: %#v", reasoning.Status) + } + if reasoning.Content != nil { + t.Fatalf("reasoning should not be converted to ordinary message content: %#v", reasoning.Content) + } + if reasoning.ResponsesReasoning == nil || len(reasoning.ResponsesReasoning.Summary) != 1 { + t.Fatalf("reasoning summary mismatch: %#v", reasoning.ResponsesReasoning) + } + summary := reasoning.ResponsesReasoning.Summary[0] + if summary.Type != schemas.ResponsesReasoningContentBlockTypeSummaryText || summary.Text != "...reasoning text..." { + t.Fatalf("reasoning summary block mismatch: %#v", summary) + } + + message := converted.Output[1] + if message.Type == nil || *message.Type != schemas.ResponsesMessageTypeMessage { + t.Fatalf("assistant output type mismatch: %#v", message.Type) + } + if message.Role == nil || *message.Role != schemas.ResponsesInputMessageRoleAssistant { + t.Fatalf("assistant output role mismatch: %#v", message.Role) + } + if message.Content == nil || len(message.Content.ContentBlocks) != 1 { + t.Fatalf("assistant content mismatch: %#v", message.Content) + } + block := message.Content.ContentBlocks[0] + if block.Type != schemas.ResponsesOutputMessageContentTypeText || block.Text == nil || *block.Text != "**Столица Франции — Париж.**" { + t.Fatalf("assistant text block mismatch: %#v", block) + } +} + +func testGigaChatResponsesConverterMapsToolCall(t *testing.T) { + t.Parallel() + + response := &GigaChatResponsesResponse{ + Model: "GigaChat-2-Max", + Messages: []GigaChatResponsesMessage{{ + Role: "assistant", + MessageID: schemas.Ptr("call-message"), + ToolsStateID: schemas.Ptr("tools-state-call"), + Content: []GigaChatResponsesContentPart{{ + FunctionCall: &GigaChatResponsesFunctionCall{ + Name: "get_weather", + Arguments: map[string]interface{}{ + "city": "Moscow", + }, + }, + }}, + FinishReason: schemas.Ptr("function_call"), + }}, + } + + converted := ToBifrostResponsesResponse(schemas.GigaChat, response) + if converted == nil { + t.Fatal("expected response, got nil") + } + if converted.Status == nil || *converted.Status != "completed" { + t.Fatalf("status mismatch: %#v", converted.Status) + } + if converted.StopReason == nil || *converted.StopReason != "tool_calls" { + t.Fatalf("stop reason mismatch: %#v", converted.StopReason) + } + if len(converted.Output) != 1 { + t.Fatalf("output count mismatch: got %d", len(converted.Output)) + } + output := converted.Output[0] + if output.Type == nil || *output.Type != schemas.ResponsesMessageTypeFunctionCall { + t.Fatalf("output type mismatch: %#v", output.Type) + } + if output.ResponsesToolMessage == nil || output.ResponsesToolMessage.Name == nil || *output.ResponsesToolMessage.Name != "get_weather" { + t.Fatalf("tool message mismatch: %#v", output.ResponsesToolMessage) + } + if output.ResponsesToolMessage.Arguments == nil || *output.ResponsesToolMessage.Arguments != `{"city":"Moscow"}` { + t.Fatalf("arguments mismatch: %#v", output.ResponsesToolMessage.Arguments) + } + assertGigaChatResponsesEncodedCallID(t, output.ResponsesToolMessage, "tools-state-call") +} + +func testGigaChatResponsesConverterMapsFunctionResultWithOriginatingCallID(t *testing.T) { + t.Parallel() + + response := &GigaChatResponsesResponse{ + Model: "GigaChat-2-Max", + Messages: []GigaChatResponsesMessage{ + { + Role: "assistant", + MessageID: schemas.Ptr("call-message"), + ToolsStateID: schemas.Ptr("tools-state-call"), + Content: []GigaChatResponsesContentPart{{ + FunctionCall: &GigaChatResponsesFunctionCall{ + Name: "get_weather", + Arguments: map[string]interface{}{"city": "Moscow"}, + }, + }}, + }, + { + Role: "tool", + MessageID: schemas.Ptr("result-message"), + ToolsStateID: schemas.Ptr("tools-state-call"), + Content: []GigaChatResponsesContentPart{{ + FunctionResult: &GigaChatResponsesFunctionResult{ + Name: "get_weather", + Result: map[string]interface{}{"temperature": 5}, + }, + }}, + }, + }, + } + + converted := ToBifrostResponsesResponse(schemas.GigaChat, response) + if converted == nil || len(converted.Output) != 2 { + t.Fatalf("converted output mismatch: %#v", converted) + } + toolCall := converted.Output[0] + toolResult := converted.Output[1] + if toolCall.Type == nil || *toolCall.Type != schemas.ResponsesMessageTypeFunctionCall { + t.Fatalf("tool call type mismatch: %#v", toolCall.Type) + } + if toolResult.Type == nil || *toolResult.Type != schemas.ResponsesMessageTypeFunctionCallOutput { + t.Fatalf("tool result type mismatch: %#v", toolResult.Type) + } + if toolCall.ResponsesToolMessage == nil || toolCall.ResponsesToolMessage.CallID == nil { + t.Fatalf("tool call id missing: %#v", toolCall.ResponsesToolMessage) + } + if toolResult.ResponsesToolMessage == nil || toolResult.ResponsesToolMessage.CallID == nil { + t.Fatalf("tool result id missing: %#v", toolResult.ResponsesToolMessage) + } + if *toolResult.ResponsesToolMessage.CallID != *toolCall.ResponsesToolMessage.CallID { + t.Fatalf("tool result call_id should match function call: call=%q result=%q", *toolCall.ResponsesToolMessage.CallID, *toolResult.ResponsesToolMessage.CallID) + } + assertGigaChatResponsesEncodedCallID(t, toolCall.ResponsesToolMessage, "tools-state-call") +} + +func testGigaChatResponsesConverterUsesUniqueCallIDsUnderSharedToolsStateID(t *testing.T) { + t.Parallel() + + response := &GigaChatResponsesResponse{ + Model: "GigaChat-2-Max", + Messages: []GigaChatResponsesMessage{ + { + Role: "assistant", + MessageID: schemas.Ptr("tool-message-1"), + ToolsStateID: schemas.Ptr("shared-tools-state"), + Content: []GigaChatResponsesContentPart{{ + FunctionCall: &GigaChatResponsesFunctionCall{ + Name: "get_weather", + Arguments: map[string]interface{}{"city": "Moscow"}, + }, + }}, + }, + { + Role: "assistant", + MessageID: schemas.Ptr("tool-message-2"), + ToolsStateID: schemas.Ptr("shared-tools-state"), + Content: []GigaChatResponsesContentPart{{ + FunctionCall: &GigaChatResponsesFunctionCall{ + Name: "get_time", + Arguments: map[string]interface{}{"city": "Moscow"}, + }, + }}, + }, + }, + } + + converted := ToBifrostResponsesResponse(schemas.GigaChat, response) + if converted == nil || len(converted.Output) != 2 { + t.Fatalf("converted output mismatch: %#v", converted) + } + firstCall := converted.Output[0].ResponsesToolMessage + secondCall := converted.Output[1].ResponsesToolMessage + firstCallID := assertGigaChatResponsesEncodedCallID(t, firstCall, "shared-tools-state") + secondCallID := assertGigaChatResponsesEncodedCallID(t, secondCall, "shared-tools-state") + if firstCallID == secondCallID { + t.Fatalf("call ids must be unique: first=%q second=%q", firstCallID, secondCallID) + } + + weatherOutput := `{"temperature":5}` + timeOutput := `{"time":"12:00"}` + request := &schemas.BifrostResponsesRequest{ + Model: "GigaChat-2-Max", + Input: []schemas.ResponsesMessage{ + converted.Output[0], + converted.Output[1], + { + Type: schemas.Ptr(schemas.ResponsesMessageTypeFunctionCallOutput), + ResponsesToolMessage: &schemas.ResponsesToolMessage{ + CallID: firstCall.CallID, + Output: &schemas.ResponsesToolMessageOutputStruct{ + ResponsesToolCallOutputStr: &weatherOutput, + }, + }, + }, + { + Type: schemas.Ptr(schemas.ResponsesMessageTypeFunctionCallOutput), + ResponsesToolMessage: &schemas.ResponsesToolMessage{ + CallID: secondCall.CallID, + Output: &schemas.ResponsesToolMessageOutputStruct{ + ResponsesToolCallOutputStr: &timeOutput, + }, + }, + }, + }, + } + + gigaChatReq, err := ToGigaChatResponsesRequest(request) + if err != nil { + t.Fatalf("ToGigaChatResponsesRequest returned error: %v", err) + } + if len(gigaChatReq.Messages) != 4 { + t.Fatalf("message count mismatch: got %d", len(gigaChatReq.Messages)) + } + assertGigaChatResponsesToolStateID(t, gigaChatReq.Messages[0], "shared-tools-state") + assertGigaChatResponsesToolStateID(t, gigaChatReq.Messages[1], "shared-tools-state") + assertGigaChatResponsesToolStateID(t, gigaChatReq.Messages[2], "shared-tools-state") + assertGigaChatResponsesToolStateID(t, gigaChatReq.Messages[3], "shared-tools-state") + if gigaChatReq.Messages[2].Content[0].FunctionResult == nil || gigaChatReq.Messages[2].Content[0].FunctionResult.Name != "get_weather" { + t.Fatalf("first function result name mismatch: %#v", gigaChatReq.Messages[2].Content) + } + if gigaChatReq.Messages[3].Content[0].FunctionResult == nil || gigaChatReq.Messages[3].Content[0].FunctionResult.Name != "get_time" { + t.Fatalf("second function result name mismatch: %#v", gigaChatReq.Messages[3].Content) + } +} + +func assertGigaChatResponsesToolStateID(t *testing.T, message GigaChatResponsesMessage, want string) { + t.Helper() + + if message.ToolsStateID == nil || *message.ToolsStateID != want { + t.Fatalf("tools_state_id mismatch: got %#v, want %q", message.ToolsStateID, want) + } +} + +func assertGigaChatResponsesEncodedCallID(t *testing.T, toolMessage *schemas.ResponsesToolMessage, toolsStateID string) string { + t.Helper() + + if toolMessage == nil || toolMessage.CallID == nil { + t.Fatalf("call id missing: %#v", toolMessage) + } + callID := strings.TrimSpace(*toolMessage.CallID) + if callID == "" { + t.Fatalf("call id is empty: %#v", toolMessage.CallID) + } + if !strings.HasPrefix(callID, gigaChatResponsesGeneratedCallIDPrefix) { + t.Fatalf("call id should use generated prefix: got %q", callID) + } + if callID == toolsStateID { + t.Fatalf("call id should not be the raw tools_state_id: got %q", callID) + } + if decoded := toGigaChatResponsesToolsStateIDFromCallID(callID); decoded != toolsStateID { + t.Fatalf("decoded tools_state_id mismatch: got %q, want %q", decoded, toolsStateID) + } + if toolMessage.Name != nil { + wantName := strings.TrimSpace(*toolMessage.Name) + if decodedName := toGigaChatResponsesFunctionNameFromCallID(callID); decodedName != wantName { + t.Fatalf("decoded function name mismatch: got %q, want %q", decodedName, wantName) + } + } + return callID +} + +func testGigaChatResponsesConverterUsesToolStateIDAliasAsCallID(t *testing.T) { + t.Parallel() + + var response GigaChatResponsesResponse + if err := json.Unmarshal([]byte(`{ + "model": "GigaChat-3-Ultra", + "messages": [{ + "role": "assistant", + "message_id": "call-message", + "tool_state_id": "019e8282-bb13-73fc-bbe8-5f52856d166b", + "content": [{ + "function_call": { + "name": "get_weather", + "arguments": {"city": "Moscow"} + } + }] + }] + }`), &response); err != nil { + t.Fatalf("failed to unmarshal response: %v", err) + } + + converted := ToBifrostResponsesResponse(schemas.GigaChat, &response) + if converted == nil || len(converted.Output) != 1 { + t.Fatalf("converted output mismatch: %#v", converted) + } + output := converted.Output[0] + if output.Type == nil || *output.Type != schemas.ResponsesMessageTypeFunctionCall { + t.Fatalf("output type mismatch: %#v", output.Type) + } + assertGigaChatResponsesEncodedCallID(t, output.ResponsesToolMessage, "019e8282-bb13-73fc-bbe8-5f52856d166b") +} + +func testGigaChatResponsesConverterFallsBackToResponseToolsStateID(t *testing.T) { + t.Parallel() + + response := &GigaChatResponsesResponse{ + Model: "GigaChat-2-Max", + ToolsStateID: schemas.Ptr("response-tools-state"), + Messages: []GigaChatResponsesMessage{ + { + Role: "assistant", + Content: []GigaChatResponsesContentPart{{ + FunctionCall: &GigaChatResponsesFunctionCall{ + Name: "get_weather", + Arguments: map[string]interface{}{"city": "Moscow"}, + }, + }}, + }, + { + Role: "assistant", + ToolsStateID: schemas.Ptr("message-tools-state"), + Content: []GigaChatResponsesContentPart{{ + FunctionCall: &GigaChatResponsesFunctionCall{ + Name: "get_time", + Arguments: map[string]interface{}{"city": "Moscow"}, + }, + }}, + }, + }, + } + + converted := ToBifrostResponsesResponse(schemas.GigaChat, response) + if converted == nil || len(converted.Output) != 2 { + t.Fatalf("converted output mismatch: %#v", converted) + } + firstCall := converted.Output[0].ResponsesToolMessage + assertGigaChatResponsesEncodedCallID(t, firstCall, "response-tools-state") + secondCall := converted.Output[1].ResponsesToolMessage + assertGigaChatResponsesEncodedCallID(t, secondCall, "message-tools-state") +} + +func testGigaChatResponsesConverterPreservesOrdinaryMessageToolStateID(t *testing.T) { + t.Parallel() + + response := &GigaChatResponsesResponse{ + Model: "GigaChat-3-Ultra", + Messages: []GigaChatResponsesMessage{{ + Role: "assistant", + MessageID: schemas.Ptr("ordinary-message"), + ToolStateID: schemas.Ptr("019e8282-bb13-73fc-bbe8-5f52856d166b"), + Content: []GigaChatResponsesContentPart{{ + Text: schemas.Ptr("Forecast: Next Tuesday brings a useful introduction."), + }}, + }}, + } + + converted := ToBifrostResponsesResponse(schemas.GigaChat, response) + if converted == nil || len(converted.Output) != 1 { + t.Fatalf("converted output mismatch: %#v", converted) + } + output := converted.Output[0] + if output.Type == nil || *output.Type != schemas.ResponsesMessageTypeMessage { + t.Fatalf("ordinary assistant output type mismatch: %#v", output.Type) + } + if output.ResponsesToolMessage != nil { + t.Fatalf("ordinary assistant message should not get tool call fields: %#v", output.ResponsesToolMessage) + } + rawStateIDs, ok := converted.ProviderExtraFields["message_tools_state_ids"].([]map[string]interface{}) + if !ok || len(rawStateIDs) != 1 { + t.Fatalf("message tool state metadata mismatch: %#v", converted.ProviderExtraFields) + } + stateID := rawStateIDs[0] + if stateID["tools_state_id"] != "019e8282-bb13-73fc-bbe8-5f52856d166b" || stateID["message_id"] != "ordinary-message" || stateID["role"] != "assistant" || stateID["index"] != 0 { + t.Fatalf("message tool state metadata mismatch: %#v", stateID) + } +} + +func testGigaChatResponsesConverterMapsThreadStorage(t *testing.T) { + t.Parallel() + + response := &GigaChatResponsesResponse{ + ThreadID: schemas.Ptr("thread-123"), + MessageID: schemas.Ptr("message-456"), + Model: "GigaChat-3-Ultra", + Messages: []GigaChatResponsesMessage{{ + Role: "assistant", + Content: []GigaChatResponsesContentPart{{ + Text: schemas.Ptr("Stored context response."), + }}, + FinishReason: schemas.Ptr("stop"), + }}, + } + + converted := ToBifrostResponsesResponse(schemas.GigaChat, response) + if converted == nil { + t.Fatal("expected response, got nil") + } + if converted.ID == nil || *converted.ID != "thread-123" { + t.Fatalf("response id should fall back to thread id, got %#v", converted.ID) + } + if converted.Conversation == nil || + converted.Conversation.ResponsesResponseConversationStruct == nil || + converted.Conversation.ResponsesResponseConversationStruct.ID != "thread-123" { + t.Fatalf("conversation mismatch: %#v", converted.Conversation) + } + if converted.ProviderExtraFields["thread_id"] != "thread-123" || converted.ProviderExtraFields["message_id"] != "message-456" { + t.Fatalf("provider extra fields mismatch: %#v", converted.ProviderExtraFields) + } +} + +func testGigaChatResponsesExecutesWithOAuthToken(t *testing.T) { + t.Parallel() + + var tokenRequests atomic.Int32 + var responsesRequests atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) { + switch request.URL.Path { + case "/oauth": + tokenRequests.Add(1) + if got := request.Header.Get("Authorization"); got != "Basic super-secret-credentials" { + t.Fatalf("token authorization header mismatch: got %q", got) + } + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"access_token":"responses-access-token","expires_at":1893456000}`)) + case "/v2/chat/completions": + responsesRequests.Add(1) + if got := request.Header.Get("Authorization"); got != "Bearer responses-access-token" { + t.Fatalf("responses authorization header mismatch: got %q", got) + } + if strings.Contains(request.Header.Get("Authorization"), "super-secret-credentials") { + t.Fatal("responses request leaked OAuth credentials") + } + assertGigaChatResponsesRequestBody(t, request) + w.Header().Set("Content-Type", "application/json") + w.Header().Set("X-Request-ID", "responses-request-id") + _, _ = w.Write([]byte(`{ + "message_id":"resp-test", + "messages":[{"role":"assistant","message_id":"msg-test","content":[{"text":"Здравствуйте"}],"finish_reason":"stop"}], + "created_at":1700000000, + "model":"GigaChat-2", + "usage":{"input_tokens":7,"input_tokens_details":{"cached_tokens":2},"output_tokens":3,"total_tokens":10} + }`)) + default: + t.Fatalf("unexpected path: %s", request.URL.Path) + } + })) + defer server.Close() + + provider := newTestGigaChatChatProvider(t, server.URL) + provider.sendBackRawRequest = true + provider.sendBackRawResponse = true + response, bifrostErr := provider.Responses(testBifrostContext(), testGigaChatOAuthKey(server.URL+"/oauth", "", "super-secret-credentials"), testGigaChatResponsesExecutionRequest()) + if bifrostErr != nil { + t.Fatalf("Responses returned error: %v", bifrostErr) + } + if tokenRequests.Load() != 1 { + t.Fatalf("token request count mismatch: got %d, want 1", tokenRequests.Load()) + } + if responsesRequests.Load() != 1 { + t.Fatalf("responses request count mismatch: got %d, want 1", responsesRequests.Load()) + } + if response.ID == nil || *response.ID != "resp-test" { + t.Fatalf("id mismatch: %#v", response.ID) + } + if response.ExtraFields.Provider != schemas.GigaChat { + t.Fatalf("provider mismatch: got %q, want %q", response.ExtraFields.Provider, schemas.GigaChat) + } + if response.Usage == nil || response.Usage.TotalTokens != 10 { + t.Fatalf("usage mismatch: %#v", response.Usage) + } + if len(response.Output) != 1 || response.Output[0].Content == nil || len(response.Output[0].Content.ContentBlocks) != 1 { + t.Fatalf("output mismatch: %#v", response.Output) + } + if got := *response.Output[0].Content.ContentBlocks[0].Text; got != "Здравствуйте" { + t.Fatalf("content mismatch: got %q", got) + } + if response.ExtraFields.RawRequest == nil || response.ExtraFields.RawResponse == nil { + t.Fatalf("expected raw request and response, got request=%#v response=%#v", response.ExtraFields.RawRequest, response.ExtraFields.RawResponse) + } + if got := response.ExtraFields.ProviderResponseHeaders["X-Request-Id"]; got != "responses-request-id" { + t.Fatalf("provider response header mismatch: %#v", response.ExtraFields.ProviderResponseHeaders) + } +} + +func testGigaChatResponsesUploadsInputImageAttachment(t *testing.T) { + t.Parallel() + + var uploadRequests atomic.Int32 + var responsesRequests atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) { + switch request.URL.Path { + case "/v1/files": + uploadRequests.Add(1) + if got := request.Header.Get("Authorization"); got != "Bearer image-token" { + t.Fatalf("file upload authorization header mismatch: got %q", got) + } + if err := request.ParseMultipartForm(1024); err != nil { + t.Fatalf("failed to parse upload multipart form: %v", err) + } + if got := request.FormValue("purpose"); got != "general" { + t.Fatalf("upload purpose mismatch: got %q", got) + } + file, header, err := request.FormFile("file") + if err != nil { + t.Fatalf("failed to read uploaded file: %v", err) + } + defer file.Close() + fileBytes, err := io.ReadAll(file) + if err != nil { + t.Fatalf("failed to read uploaded bytes: %v", err) + } + if string(fileBytes) != "image-bytes" { + t.Fatalf("uploaded image bytes mismatch: %q", fileBytes) + } + if header.Filename != "image.png" { + t.Fatalf("uploaded image filename mismatch: got %q", header.Filename) + } + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"id":"uploaded-image","object":"file","bytes":11,"created_at":1700000000,"filename":"image.png","purpose":"general"}`)) + case "/v2/chat/completions": + responsesRequests.Add(1) + body, err := io.ReadAll(request.Body) + if err != nil { + t.Fatalf("failed to read responses body: %v", err) + } + assertGigaChatResponsesBodyFile(t, body, "uploaded-image") + bodyStr := string(body) + if strings.Contains(bodyStr, "data:image") || strings.Contains(bodyStr, "image_url") { + t.Fatalf("responses body leaked OpenAI image payload: %s", body) + } + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"messages":[{"role":"assistant","content":[{"text":"На изображении..."}],"finish_reason":"stop"}],"model":"GigaChat-2"}`)) + default: + t.Fatalf("unexpected path: %s", request.URL.Path) + } + })) + defer server.Close() + + prompt := "What's in this image?" + imageURL := "data:image/png;base64,aW1hZ2UtYnl0ZXM=" + request := &schemas.BifrostResponsesRequest{ + Model: "GigaChat-2", + Input: []schemas.ResponsesMessage{{ + Role: schemas.Ptr(schemas.ResponsesInputMessageRoleUser), + Content: &schemas.ResponsesMessageContent{ContentBlocks: []schemas.ResponsesMessageContentBlock{ + {Type: schemas.ResponsesInputMessageContentBlockTypeText, Text: &prompt}, + { + Type: schemas.ResponsesInputMessageContentBlockTypeImage, + ResponsesInputMessageContentBlockImage: &schemas.ResponsesInputMessageContentBlockImage{ + ImageURL: &imageURL, + }, + }, + }}, + }}, + } + + provider := newTestGigaChatChatProvider(t, server.URL) + response, bifrostErr := provider.Responses(testBifrostContext(), testGigaChatAccessTokenKey("image-token"), request) + if bifrostErr != nil { + t.Fatalf("Responses returned error: %v", bifrostErr) + } + if response == nil { + t.Fatal("expected response, got nil") + } + if uploadRequests.Load() != 1 { + t.Fatalf("upload request count mismatch: got %d, want 1", uploadRequests.Load()) + } + if responsesRequests.Load() != 1 { + t.Fatalf("responses request count mismatch: got %d, want 1", responsesRequests.Load()) + } +} + +func testGigaChatResponsesUploadsInlineFileAttachment(t *testing.T) { + t.Parallel() + + var uploadRequests atomic.Int32 + var responsesRequests atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) { + switch request.URL.Path { + case "/v1/files": + uploadRequests.Add(1) + if err := request.ParseMultipartForm(1024); err != nil { + t.Fatalf("failed to parse upload multipart form: %v", err) + } + file, header, err := request.FormFile("file") + if err != nil { + t.Fatalf("failed to read uploaded file: %v", err) + } + defer file.Close() + fileBytes, err := io.ReadAll(file) + if err != nil { + t.Fatalf("failed to read uploaded bytes: %v", err) + } + if string(fileBytes) != "%PDF test" { + t.Fatalf("uploaded file bytes mismatch: %q", fileBytes) + } + if header.Filename != "report.pdf" { + t.Fatalf("uploaded filename mismatch: got %q", header.Filename) + } + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"id":"uploaded-pdf","object":"file","bytes":9,"created_at":1700000000,"filename":"report.pdf","purpose":"general"}`)) + case "/v2/chat/completions": + responsesRequests.Add(1) + body, err := io.ReadAll(request.Body) + if err != nil { + t.Fatalf("failed to read responses body: %v", err) + } + assertGigaChatResponsesBodyFile(t, body, "uploaded-pdf") + bodyStr := string(body) + if strings.Contains(bodyStr, "file_data") || strings.Contains(bodyStr, "application/pdf;base64") { + t.Fatalf("responses body leaked OpenAI file payload: %s", body) + } + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"messages":[{"role":"assistant","content":[{"text":"Краткое содержание..."}],"finish_reason":"stop"}],"model":"GigaChat-2"}`)) + default: + t.Fatalf("unexpected path: %s", request.URL.Path) + } + })) + defer server.Close() + + prompt := "Summarize this file." + filename := "report.pdf" + fileType := "application/pdf" + fileData := "data:application/pdf;base64,JVBERiB0ZXN0" + request := &schemas.BifrostResponsesRequest{ + Model: "GigaChat-2", + Input: []schemas.ResponsesMessage{{ + Role: schemas.Ptr(schemas.ResponsesInputMessageRoleUser), + Content: &schemas.ResponsesMessageContent{ContentBlocks: []schemas.ResponsesMessageContentBlock{ + {Type: schemas.ResponsesInputMessageContentBlockTypeText, Text: &prompt}, + { + Type: schemas.ResponsesInputMessageContentBlockTypeFile, + ResponsesInputMessageContentBlockFile: &schemas.ResponsesInputMessageContentBlockFile{ + Filename: &filename, + FileType: &fileType, + FileData: &fileData, + }, + }, + }}, + }}, + } + + provider := newTestGigaChatChatProvider(t, server.URL) + response, bifrostErr := provider.Responses(testBifrostContext(), testGigaChatAccessTokenKey("file-token"), request) + if bifrostErr != nil { + t.Fatalf("Responses returned error: %v", bifrostErr) + } + if response == nil { + t.Fatal("expected response, got nil") + } + if uploadRequests.Load() != 1 { + t.Fatalf("upload request count mismatch: got %d, want 1", uploadRequests.Load()) + } + if responsesRequests.Load() != 1 { + t.Fatalf("responses request count mismatch: got %d, want 1", responsesRequests.Load()) + } +} + +func testGigaChatResponsesReusesUploadedAttachmentAfterBackendError(t *testing.T) { + t.Parallel() + + var uploadRequests atomic.Int32 + var responsesRequests atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) { + switch request.URL.Path { + case "/v1/files": + uploadRequests.Add(1) + if err := request.ParseMultipartForm(1024); err != nil { + t.Fatalf("failed to parse upload multipart form: %v", err) + } + file, _, err := request.FormFile("file") + if err != nil { + t.Fatalf("failed to read uploaded file: %v", err) + } + defer file.Close() + fileBytes, err := io.ReadAll(file) + if err != nil { + t.Fatalf("failed to read uploaded bytes: %v", err) + } + if string(fileBytes) != "%PDF retry" { + t.Fatalf("uploaded file bytes mismatch: %q", fileBytes) + } + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"id":"uploaded-retry-pdf","object":"file","bytes":10,"created_at":1700000000,"filename":"retry.pdf","purpose":"general"}`)) + case "/v2/chat/completions": + requestIndex := responsesRequests.Add(1) + body, err := io.ReadAll(request.Body) + if err != nil { + t.Fatalf("failed to read responses body: %v", err) + } + assertGigaChatResponsesBodyFile(t, body, "uploaded-retry-pdf") + if requestIndex == 1 { + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusInternalServerError) + _, _ = w.Write([]byte(`{"status":500,"message":"temporary backend failure"}`)) + return + } + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"messages":[{"role":"assistant","content":[{"text":"ok"}],"finish_reason":"stop"}],"model":"GigaChat-2"}`)) + default: + t.Fatalf("unexpected path: %s", request.URL.Path) + } + })) + defer server.Close() + + prompt := "Summarize this file." + filename := "retry.pdf" + fileType := "application/pdf" + fileData := "data:application/pdf;base64,JVBERiByZXRyeQ==" + request := &schemas.BifrostResponsesRequest{ + Model: "GigaChat-2", + Input: []schemas.ResponsesMessage{{ + Role: schemas.Ptr(schemas.ResponsesInputMessageRoleUser), + Content: &schemas.ResponsesMessageContent{ContentBlocks: []schemas.ResponsesMessageContentBlock{ + {Type: schemas.ResponsesInputMessageContentBlockTypeText, Text: &prompt}, + { + Type: schemas.ResponsesInputMessageContentBlockTypeFile, + ResponsesInputMessageContentBlockFile: &schemas.ResponsesInputMessageContentBlockFile{ + Filename: &filename, + FileType: &fileType, + FileData: &fileData, + }, + }, + }}, + }}, + } + + provider := newTestGigaChatChatProvider(t, server.URL) + ctx := testBifrostContext() + key := testGigaChatAccessTokenKey("file-token") + + firstResponse, firstErr := provider.Responses(ctx, key, request) + if firstResponse != nil { + t.Fatalf("expected nil response from first backend failure, got %#v", firstResponse) + } + if firstErr == nil || firstErr.StatusCode == nil || *firstErr.StatusCode != http.StatusInternalServerError { + t.Fatalf("expected first backend 500, got %#v", firstErr) + } + + response, bifrostErr := provider.Responses(ctx, key, request) + if bifrostErr != nil { + t.Fatalf("second Responses returned error: %v", bifrostErr) + } + if response == nil { + t.Fatal("expected second response, got nil") + } + if uploadRequests.Load() != 1 { + t.Fatalf("upload request count mismatch: got %d, want 1", uploadRequests.Load()) + } + if responsesRequests.Load() != 2 { + t.Fatalf("responses request count mismatch: got %d, want 2", responsesRequests.Load()) + } +} + +func testGigaChatResponsesReusesCompletedUploadsAfterPartialAttachmentFailure(t *testing.T) { + t.Parallel() + + var firstFileUploads atomic.Int32 + var secondFileUploads atomic.Int32 + var responsesRequests atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) { + switch request.URL.Path { + case "/v1/files": + if err := request.ParseMultipartForm(1024); err != nil { + t.Fatalf("failed to parse upload multipart form: %v", err) + } + file, _, err := request.FormFile("file") + if err != nil { + t.Fatalf("failed to read uploaded file: %v", err) + } + defer file.Close() + fileBytes, err := io.ReadAll(file) + if err != nil { + t.Fatalf("failed to read uploaded bytes: %v", err) + } + + w.Header().Set("Content-Type", "application/json") + switch string(fileBytes) { + case "%PDF first": + firstFileUploads.Add(1) + _, _ = w.Write([]byte(`{"id":"uploaded-first-pdf","object":"file","bytes":10,"created_at":1700000000,"filename":"first.pdf","purpose":"general"}`)) + case "%PDF second": + if secondFileUploads.Add(1) == 1 { + w.WriteHeader(http.StatusInternalServerError) + _, _ = w.Write([]byte(`{"status":500,"message":"temporary second upload failure"}`)) + return + } + _, _ = w.Write([]byte(`{"id":"uploaded-second-pdf","object":"file","bytes":11,"created_at":1700000000,"filename":"second.pdf","purpose":"general"}`)) + default: + t.Fatalf("unexpected uploaded file bytes: %q", fileBytes) + } + case "/v2/chat/completions": + responsesRequests.Add(1) + body, err := io.ReadAll(request.Body) + if err != nil { + t.Fatalf("failed to read responses body: %v", err) + } + var payload struct { + Messages []struct { + Content []struct { + Files []struct { + ID string `json:"id"` + } `json:"files"` + } `json:"content"` + } `json:"messages"` + } + if err := json.Unmarshal(body, &payload); err != nil { + t.Fatalf("failed to unmarshal responses body %s: %v", body, err) + } + if len(payload.Messages) != 1 || len(payload.Messages[0].Content) != 3 { + t.Fatalf("responses body content mismatch: %s", body) + } + firstFiles := payload.Messages[0].Content[1].Files + secondFiles := payload.Messages[0].Content[2].Files + if len(firstFiles) != 1 || firstFiles[0].ID != "uploaded-first-pdf" { + t.Fatalf("first uploaded file id mismatch: %#v body %s", firstFiles, body) + } + if len(secondFiles) != 1 || secondFiles[0].ID != "uploaded-second-pdf" { + t.Fatalf("second uploaded file id mismatch: %#v body %s", secondFiles, body) + } + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"messages":[{"role":"assistant","content":[{"text":"ok"}],"finish_reason":"stop"}],"model":"GigaChat-2"}`)) + default: + t.Fatalf("unexpected path: %s", request.URL.Path) + } + })) + defer server.Close() + + prompt := "Compare these files." + firstFilename := "first.pdf" + secondFilename := "second.pdf" + fileType := "application/pdf" + firstFileData := "data:application/pdf;base64,JVBERiBmaXJzdA==" + secondFileData := "data:application/pdf;base64,JVBERiBzZWNvbmQ=" + request := &schemas.BifrostResponsesRequest{ + Model: "GigaChat-2", + Input: []schemas.ResponsesMessage{{ + Role: schemas.Ptr(schemas.ResponsesInputMessageRoleUser), + Content: &schemas.ResponsesMessageContent{ContentBlocks: []schemas.ResponsesMessageContentBlock{ + {Type: schemas.ResponsesInputMessageContentBlockTypeText, Text: &prompt}, + { + Type: schemas.ResponsesInputMessageContentBlockTypeFile, + ResponsesInputMessageContentBlockFile: &schemas.ResponsesInputMessageContentBlockFile{ + Filename: &firstFilename, + FileType: &fileType, + FileData: &firstFileData, + }, + }, + { + Type: schemas.ResponsesInputMessageContentBlockTypeFile, + ResponsesInputMessageContentBlockFile: &schemas.ResponsesInputMessageContentBlockFile{ + Filename: &secondFilename, + FileType: &fileType, + FileData: &secondFileData, + }, + }, + }}, + }}, + } + + provider := newTestGigaChatChatProvider(t, server.URL) + ctx := testBifrostContext() + key := testGigaChatAccessTokenKey("file-token") + + firstResponse, firstErr := provider.Responses(ctx, key, request) + if firstResponse != nil { + t.Fatalf("expected nil response from partial upload failure, got %#v", firstResponse) + } + if firstErr == nil || firstErr.StatusCode == nil || *firstErr.StatusCode != http.StatusInternalServerError { + t.Fatalf("expected second attachment upload 500, got %#v", firstErr) + } + + response, bifrostErr := provider.Responses(ctx, key, request) + if bifrostErr != nil { + t.Fatalf("second Responses returned error: %v", bifrostErr) + } + if response == nil { + t.Fatal("expected second response, got nil") + } + if firstFileUploads.Load() != 1 { + t.Fatalf("first file upload count mismatch: got %d, want 1", firstFileUploads.Load()) + } + if secondFileUploads.Load() != 2 { + t.Fatalf("second file upload count mismatch: got %d, want 2", secondFileUploads.Load()) + } + if responsesRequests.Load() != 1 { + t.Fatalf("responses request count mismatch: got %d, want 1", responsesRequests.Load()) + } +} + +func testGigaChatResponsesMapsProviderErrors(t *testing.T) { + t.Parallel() + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusBadRequest) + _, _ = w.Write([]byte(`{"status":400,"code":123,"message":"bad responses request"}`)) + })) + defer server.Close() + + provider := newTestGigaChatChatProvider(t, server.URL) + response, bifrostErr := provider.Responses(testBifrostContext(), testGigaChatAccessTokenKey("provider-error-token"), testGigaChatResponsesExecutionRequest()) + if response != nil { + t.Fatalf("expected nil response, got %#v", response) + } + if bifrostErr == nil { + t.Fatal("expected provider error, got nil") + } + if bifrostErr.StatusCode == nil || *bifrostErr.StatusCode != http.StatusBadRequest { + t.Fatalf("status mismatch: %#v", bifrostErr.StatusCode) + } + if bifrostErr.Error == nil || bifrostErr.Error.Message != "bad responses request" { + t.Fatalf("message mismatch: %#v", bifrostErr.Error) + } + if bifrostErr.Error.Code == nil || *bifrostErr.Error.Code != "123" { + t.Fatalf("code mismatch: %#v", bifrostErr.Error) + } + assertNoGigaChatSecretLeak(t, bifrostErr.String()) +} + +func testGigaChatResponsesRefreshesTokenAfterUnauthorized(t *testing.T) { + t.Parallel() + + var tokenRequests atomic.Int32 + var responsesRequests atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) { + switch request.URL.Path { + case "/oauth": + tokenIndex := tokenRequests.Add(1) + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(fmt.Sprintf(`{"access_token":"responses-token-%d","expires_at":1893456000}`, tokenIndex))) + case "/v2/chat/completions": + responsesIndex := responsesRequests.Add(1) + wantAuthorization := fmt.Sprintf("Bearer responses-token-%d", responsesIndex) + if got := request.Header.Get("Authorization"); got != wantAuthorization { + t.Fatalf("authorization header mismatch on request %d: got %q, want %q", responsesIndex, got, wantAuthorization) + } + w.Header().Set("Content-Type", "application/json") + if responsesIndex == 1 { + w.WriteHeader(http.StatusUnauthorized) + _, _ = w.Write([]byte(`{"status":401,"message":"expired token"}`)) + return + } + _, _ = w.Write([]byte(`{"messages":[{"role":"assistant","content":[{"text":"ok"}],"finish_reason":"stop"}],"model":"GigaChat-2"}`)) + default: + t.Fatalf("unexpected path: %s", request.URL.Path) + } + })) + defer server.Close() + + provider := newTestGigaChatChatProvider(t, server.URL) + response, bifrostErr := provider.Responses(testBifrostContext(), testGigaChatOAuthKey(server.URL+"/oauth", "", "test-credentials"), testGigaChatResponsesExecutionRequest()) + if bifrostErr != nil { + t.Fatalf("Responses returned error: %v", bifrostErr) + } + if response == nil || len(response.Output) != 1 { + t.Fatalf("unexpected response: %#v", response) + } + if tokenRequests.Load() != 2 { + t.Fatalf("token request count mismatch: got %d, want 2", tokenRequests.Load()) + } + if responsesRequests.Load() != 2 { + t.Fatalf("responses request count mismatch: got %d, want 2", responsesRequests.Load()) + } +} + +func testGigaChatResponsesStreamTextDeltasAndUsage(t *testing.T) { + t.Parallel() + + var tokenRequests atomic.Int32 + var streamRequests atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) { + switch request.URL.Path { + case "/oauth": + tokenRequests.Add(1) + if got := request.Header.Get("Authorization"); got != "Basic super-secret-credentials" { + t.Fatalf("token authorization header mismatch: got %q", got) + } + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"access_token":"responses-stream-token","expires_at":1893456000}`)) + case "/v2/chat/completions": + streamRequests.Add(1) + if got := request.Header.Get("Authorization"); got != "Bearer responses-stream-token" { + t.Fatalf("stream authorization header mismatch: got %q", got) + } + if strings.Contains(request.Header.Get("Authorization"), "super-secret-credentials") { + t.Fatal("stream request leaked OAuth credentials") + } + assertGigaChatResponsesStreamRequestBody(t, request) + w.Header().Set("Content-Type", "text/event-stream") + w.Header().Set("X-Request-ID", "responses-stream-request-id") + _, _ = w.Write([]byte("data: {\"event\":\"message\",\"message_id\":\"resp-stream\",\"messages\":[{\"role\":\"assistant\",\"content\":[{\"text\":\"При\"}]}],\"created_at\":1700000000,\"model\":\"GigaChat-2\"}\n\n")) + _, _ = w.Write([]byte("data: {\"event\":\"message\",\"message_id\":\"resp-stream\",\"messages\":[{\"content\":[{\"text\":\"вет\"}]}],\"created_at\":1700000000,\"model\":\"GigaChat-2\"}\n\n")) + _, _ = w.Write([]byte("data: {\"event\":\"done\",\"message_id\":\"resp-stream\",\"messages\":[{\"finish_reason\":\"stop\"}],\"created_at\":1700000000,\"model\":\"GigaChat-2\",\"usage\":{\"input_tokens\":7,\"input_tokens_details\":{\"cached_tokens\":2},\"output_tokens\":3,\"total_tokens\":10}}\n\n")) + _, _ = w.Write([]byte("data: [DONE]\n\n")) + default: + t.Fatalf("unexpected path: %s", request.URL.Path) + } + })) + defer server.Close() + + provider := newTestGigaChatChatProvider(t, server.URL) + provider.sendBackRawRequest = true + provider.sendBackRawResponse = true + ctx := testBifrostContext() + + stream, bifrostErr := provider.ResponsesStream(ctx, testGigaChatPostHookRunner, nil, testGigaChatOAuthKey(server.URL+"/oauth", "", "super-secret-credentials"), testGigaChatResponsesExecutionRequest()) + if bifrostErr != nil { + t.Fatalf("ResponsesStream returned error: %v", bifrostErr) + } + + chunks := collectGigaChatStreamChunks(t, stream) + if tokenRequests.Load() != 1 { + t.Fatalf("token request count mismatch: got %d, want 1", tokenRequests.Load()) + } + if streamRequests.Load() != 1 { + t.Fatalf("stream request count mismatch: got %d, want 1", streamRequests.Load()) + } + + responses := collectGigaChatResponsesStreamResponses(t, chunks) + assertGigaChatResponsesStreamTypes(t, responses, []schemas.ResponsesStreamResponseType{ + schemas.ResponsesStreamResponseTypeCreated, + schemas.ResponsesStreamResponseTypeInProgress, + schemas.ResponsesStreamResponseTypeOutputItemAdded, + schemas.ResponsesStreamResponseTypeContentPartAdded, + schemas.ResponsesStreamResponseTypeOutputTextDelta, + schemas.ResponsesStreamResponseTypeOutputTextDelta, + schemas.ResponsesStreamResponseTypeOutputTextDone, + schemas.ResponsesStreamResponseTypeContentPartDone, + schemas.ResponsesStreamResponseTypeOutputItemDone, + schemas.ResponsesStreamResponseTypeCompleted, + }) + if responses[4].Delta == nil || *responses[4].Delta != "При" { + t.Fatalf("first delta mismatch: %#v", responses[4].Delta) + } + if responses[5].Delta == nil || *responses[5].Delta != "вет" { + t.Fatalf("second delta mismatch: %#v", responses[5].Delta) + } + finalResponse := responses[len(responses)-1] + if finalResponse.Response == nil || finalResponse.Response.Usage == nil || finalResponse.Response.Usage.TotalTokens != 10 { + t.Fatalf("final usage mismatch: %#v", finalResponse.Response) + } + if finalResponse.Response.Usage.InputTokensDetails == nil || finalResponse.Response.Usage.InputTokensDetails.CachedReadTokens != 2 { + t.Fatalf("cached token usage mismatch: %#v", finalResponse.Response.Usage) + } + if finalResponse.Response.Status == nil || *finalResponse.Response.Status != "completed" { + t.Fatalf("final status mismatch: %#v", finalResponse.Response.Status) + } + if finalResponse.ExtraFields.RawRequest == nil || finalResponse.ExtraFields.RawResponse == nil { + t.Fatalf("expected raw request and response, got request=%#v response=%#v", finalResponse.ExtraFields.RawRequest, finalResponse.ExtraFields.RawResponse) + } + if got := ctx.Value(schemas.BifrostContextKeyProviderResponseHeaders); got == nil { + t.Fatal("provider response headers were not stored in context") + } +} + +func testGigaChatResponsesStreamReasoningDeltas(t *testing.T) { + t.Parallel() + + state := schemas.AcquireChatToResponsesStreamState() + defer schemas.ReleaseChatToResponsesStreamState(state) + + response := &GigaChatResponsesResponse{ + MessageID: schemas.Ptr("resp-reasoning-stream"), + CreatedAt: 1780306293, + Model: "GigaChat-2-Reasoning:2.0.29.05", + Messages: []GigaChatResponsesMessage{{ + Role: "reasoning", + Content: []GigaChatResponsesContentPart{{ + Text: schemas.Ptr("streamed reasoning"), + }}, + }}, + } + + events := ToBifrostResponsesStreamResponse(schemas.GigaChat, response, state) + if len(events) == 0 { + t.Fatal("expected stream events, got none") + } + + var foundReasoningDelta bool + for _, event := range events { + if event == nil { + continue + } + if event.Type == schemas.ResponsesStreamResponseTypeOutputTextDelta && event.Delta != nil && *event.Delta == "streamed reasoning" { + t.Fatalf("reasoning delta was emitted as output_text: %#v", event) + } + if event.Type == schemas.ResponsesStreamResponseTypeOutputItemAdded && event.Item != nil && event.Item.Role != nil && *event.Item.Role == schemas.ResponsesMessageRoleType("reasoning") { + t.Fatalf("reasoning delta created ordinary message role=reasoning: %#v", event.Item) + } + if event.Type == schemas.ResponsesStreamResponseTypeReasoningSummaryTextDelta && event.Delta != nil && *event.Delta == "streamed reasoning" { + foundReasoningDelta = true + } + } + if !foundReasoningDelta { + t.Fatalf("expected reasoning summary delta, got %#v", events) + } +} + +func testGigaChatResponsesStreamToolCallDeltas(t *testing.T) { + t.Parallel() + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) { + if request.URL.Path != "/v2/chat/completions" { + t.Fatalf("unexpected path: %s", request.URL.Path) + } + assertGigaChatResponsesStreamRequestBody(t, request) + w.Header().Set("Content-Type", "text/event-stream") + _, _ = w.Write([]byte("data: {\"event\":\"message\",\"message_id\":\"resp-tools\",\"messages\":[{\"role\":\"assistant\",\"tools_state_id\":\"call-weather\",\"content\":[{\"function_call\":{\"name\":\"get_weather\",\"arguments\":{\"city\":\"Moscow\"}}}]}],\"created_at\":1700000000,\"model\":\"GigaChat-2\"}\n\n")) + _, _ = w.Write([]byte("data: {\"event\":\"done\",\"message_id\":\"resp-tools\",\"messages\":[{\"finish_reason\":\"function_call\"}],\"created_at\":1700000000,\"model\":\"GigaChat-2\",\"usage\":{\"input_tokens\":11,\"output_tokens\":4,\"total_tokens\":15}}\n\n")) + _, _ = w.Write([]byte("data: [DONE]\n\n")) + })) + defer server.Close() + + provider := newTestGigaChatChatProvider(t, server.URL) + stream, bifrostErr := provider.ResponsesStream(testBifrostContext(), testGigaChatPostHookRunner, nil, testGigaChatAccessTokenKey("responses-stream-token"), testGigaChatResponsesExecutionRequest()) + if bifrostErr != nil { + t.Fatalf("ResponsesStream returned error: %v", bifrostErr) + } + + responses := collectGigaChatResponsesStreamResponses(t, collectGigaChatStreamChunks(t, stream)) + assertGigaChatResponsesStreamTypes(t, responses, []schemas.ResponsesStreamResponseType{ + schemas.ResponsesStreamResponseTypeCreated, + schemas.ResponsesStreamResponseTypeInProgress, + schemas.ResponsesStreamResponseTypeOutputItemAdded, + schemas.ResponsesStreamResponseTypeFunctionCallArgumentsDelta, + schemas.ResponsesStreamResponseTypeFunctionCallArgumentsDone, + schemas.ResponsesStreamResponseTypeOutputItemDone, + schemas.ResponsesStreamResponseTypeCompleted, + }) + if responses[2].Item == nil || responses[2].Item.ResponsesToolMessage == nil || responses[2].Item.ResponsesToolMessage.Name == nil || *responses[2].Item.ResponsesToolMessage.Name != "get_weather" { + t.Fatalf("tool item mismatch: %#v", responses[2].Item) + } + if responses[3].Delta == nil || *responses[3].Delta != `{"city":"Moscow"}` { + t.Fatalf("tool delta mismatch: %#v", responses[3].Delta) + } + if responses[4].Arguments == nil || *responses[4].Arguments != `{"city":"Moscow"}` { + t.Fatalf("tool arguments mismatch: %#v", responses[4].Arguments) + } + finalResponse := responses[len(responses)-1] + if finalResponse.Response == nil || finalResponse.Response.Usage == nil || finalResponse.Response.Usage.TotalTokens != 15 { + t.Fatalf("final usage mismatch: %#v", finalResponse.Response) + } +} + +func testGigaChatResponsesStreamClosesOnMessageDoneEvent(t *testing.T) { + t.Parallel() + + releaseServer := make(chan struct{}) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) { + if request.URL.Path != "/v2/chat/completions" { + t.Fatalf("unexpected path: %s", request.URL.Path) + } + assertGigaChatResponsesStreamRequestBody(t, request) + w.Header().Set("Content-Type", "text/event-stream") + flusher, ok := w.(http.Flusher) + if !ok { + t.Fatal("response writer does not support flushing") + } + _, _ = w.Write([]byte("event: response.message.delta\ndata: {\"message_id\":\"resp-done-event\",\"messages\":[{\"role\":\"assistant\",\"content\":[{\"text\":\"Привет\"}]}],\"created_at\":1780315871,\"model\":\"GigaChat-2-Reasoning:2.0.29.05\"}\n\n")) + flusher.Flush() + _, _ = w.Write([]byte("event: response.message.done\ndata: {\"model\":\"GigaChat-2-Reasoning:2.0.29.05\",\"created_at\":1780315871,\"finish_reason\":\"stop\",\"usage\":{\"input_tokens\":29,\"input_tokens_details\":{\"prompt_tokens\":29,\"cached_tokens\":3},\"output_tokens\":109,\"total_tokens\":138}}\n\n")) + flusher.Flush() + select { + case <-request.Context().Done(): + case <-releaseServer: + } + })) + defer server.Close() + defer close(releaseServer) + + provider := newTestGigaChatChatProvider(t, server.URL) + stream, bifrostErr := provider.ResponsesStream(testBifrostContext(), testGigaChatPostHookRunner, nil, testGigaChatAccessTokenKey("responses-stream-token"), testGigaChatResponsesExecutionRequest()) + if bifrostErr != nil { + t.Fatalf("ResponsesStream returned error: %v", bifrostErr) + } + + chunksDone := make(chan []*schemas.BifrostStreamChunk, 1) + go func() { + chunks := make([]*schemas.BifrostStreamChunk, 0) + for chunk := range stream { + chunks = append(chunks, chunk) + } + chunksDone <- chunks + }() + + var chunks []*schemas.BifrostStreamChunk + select { + case chunks = <-chunksDone: + case <-time.After(time.Second): + t.Fatal("responses stream did not close after response.message.done") + } + + responses := collectGigaChatResponsesStreamResponses(t, chunks) + assertGigaChatResponsesStreamTypes(t, responses, []schemas.ResponsesStreamResponseType{ + schemas.ResponsesStreamResponseTypeCreated, + schemas.ResponsesStreamResponseTypeInProgress, + schemas.ResponsesStreamResponseTypeOutputItemAdded, + schemas.ResponsesStreamResponseTypeContentPartAdded, + schemas.ResponsesStreamResponseTypeOutputTextDelta, + schemas.ResponsesStreamResponseTypeOutputTextDone, + schemas.ResponsesStreamResponseTypeContentPartDone, + schemas.ResponsesStreamResponseTypeOutputItemDone, + schemas.ResponsesStreamResponseTypeCompleted, + }) + finalResponse := responses[len(responses)-1] + if finalResponse.Response == nil || finalResponse.Response.Usage == nil || finalResponse.Response.Usage.TotalTokens != 138 { + t.Fatalf("final usage mismatch: %#v", finalResponse.Response) + } +} + +func testGigaChatResponsesStreamMapsErrorEvents(t *testing.T) { + t.Parallel() + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) { + if request.URL.Path != "/v2/chat/completions" { + t.Fatalf("unexpected path: %s", request.URL.Path) + } + w.Header().Set("Content-Type", "text/event-stream") + _, _ = w.Write([]byte("data: {\"status\":429,\"code\":42901,\"message\":\"rate limit\"}\n\n")) + _, _ = w.Write([]byte("data: [DONE]\n\n")) + })) + defer server.Close() + + provider := newTestGigaChatChatProvider(t, server.URL) + stream, bifrostErr := provider.ResponsesStream(testBifrostContext(), testGigaChatPostHookRunner, nil, testGigaChatAccessTokenKey("responses-stream-token"), testGigaChatResponsesExecutionRequest()) + if bifrostErr != nil { + t.Fatalf("ResponsesStream returned error before stream: %v", bifrostErr) + } + + chunks := collectGigaChatStreamChunks(t, stream) + if len(chunks) != 1 || chunks[0].BifrostError == nil { + t.Fatalf("expected one error chunk, got %#v", chunks) + } + streamErr := chunks[0].BifrostError + if streamErr.StatusCode == nil || *streamErr.StatusCode != http.StatusTooManyRequests { + t.Fatalf("status mismatch: %#v", streamErr.StatusCode) + } + if streamErr.Error == nil || streamErr.Error.Message != "rate limit" { + t.Fatalf("message mismatch: %#v", streamErr.Error) + } + if streamErr.Error.Code == nil || *streamErr.Error.Code != "42901" { + t.Fatalf("code mismatch: %#v", streamErr.Error) + } +} + +func testGigaChatResponsesStreamHandlesContextCancellation(t *testing.T) { + t.Parallel() + + firstChunkWritten := make(chan struct{}) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) { + if request.URL.Path != "/v2/chat/completions" { + t.Fatalf("unexpected path: %s", request.URL.Path) + } + w.Header().Set("Content-Type", "text/event-stream") + _, _ = w.Write([]byte("data: {\"event\":\"message\",\"message_id\":\"resp-cancel\",\"messages\":[{\"role\":\"assistant\",\"content\":[{\"text\":\"partial\"}]}],\"created_at\":1700000000,\"model\":\"GigaChat-2\"}\n\n")) + if flusher, ok := w.(http.Flusher); ok { + flusher.Flush() + } + close(firstChunkWritten) + <-request.Context().Done() + })) + defer server.Close() + + provider := newTestGigaChatChatProvider(t, server.URL) + ctx, cancel := schemas.NewBifrostContextWithCancel(context.Background()) + stream, bifrostErr := provider.ResponsesStream(ctx, testGigaChatPostHookRunner, nil, testGigaChatAccessTokenKey("responses-stream-token"), testGigaChatResponsesExecutionRequest()) + if bifrostErr != nil { + t.Fatalf("ResponsesStream returned error: %v", bifrostErr) + } + + select { + case <-firstChunkWritten: + case <-time.After(time.Second): + t.Fatal("timed out waiting for first stream chunk") + } + + select { + case firstChunk := <-stream: + if firstChunk == nil || firstChunk.BifrostResponsesStreamResponse == nil { + t.Fatalf("missing first responses stream chunk: %#v", firstChunk) + } + case <-time.After(time.Second): + t.Fatal("timed out waiting for first response chunk") + } + + cancel() + select { + case <-ctx.Done(): + case <-time.After(time.Second): + t.Fatal("timed out waiting for context cancellation") + } + + streamClosed := make(chan struct{}) + go func() { + for range stream { + } + close(streamClosed) + }() + + select { + case <-streamClosed: + case <-time.After(time.Second): + t.Fatal("timed out waiting for stream to close after context cancellation") + } +} + +func testGigaChatResponsesStreamPassthroughResponseOwnedByLargeReader(t *testing.T) { + t.Parallel() + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) { + if request.URL.Path != "/v2/chat/completions" { + t.Fatalf("unexpected path: %s", request.URL.Path) + } + assertGigaChatResponsesStreamRequestBody(t, request) + w.Header().Set("Content-Type", "text/event-stream") + _, _ = w.Write([]byte("data: {\"event\":\"message\",\"message_id\":\"resp-large\",\"messages\":[{\"role\":\"assistant\",\"content\":[{\"text\":\"large\"}]}],\"created_at\":1700000000,\"model\":\"GigaChat-2\"}\n\n")) + _, _ = w.Write([]byte("data: [DONE]\n\n")) + })) + defer server.Close() + + provider := newTestGigaChatChatProvider(t, server.URL) + ctx := testBifrostContext() + ctx.SetValue(schemas.BifrostContextKeyLargePayloadMode, true) + + var finalizerCalls atomic.Int32 + stream, bifrostErr := provider.ResponsesStream(ctx, testGigaChatPostHookRunner, func(context.Context) { + finalizerCalls.Add(1) + }, testGigaChatAccessTokenKey("responses-stream-token"), testGigaChatResponsesExecutionRequest()) + if bifrostErr != nil { + t.Fatalf("ResponsesStream returned error: %v", bifrostErr) + } + + select { + case chunk, ok := <-stream: + if ok { + t.Fatalf("passthrough stream channel should be closed without chunks, got %#v", chunk) + } + case <-time.After(time.Second): + t.Fatal("timed out waiting for passthrough stream channel to close") + } + if got := finalizerCalls.Load(); got != 0 { + t.Fatalf("finalizer ran before passthrough delivery: got %d calls", got) + } + + reader, ok := ctx.Value(schemas.BifrostContextKeyLargeResponseReader).(io.ReadCloser) + if !ok || reader == nil { + t.Fatalf("large response reader missing from context: %#v", ctx.Value(schemas.BifrostContextKeyLargeResponseReader)) + } + finalizingReader, ok := reader.(*gigaChatPassthroughReadCloser) + if !ok { + t.Fatalf("passthrough reader type mismatch: %T", reader) + } + largeReader, ok := finalizingReader.ReadCloser.(*providerUtils.LargeResponseReader) + if !ok { + t.Fatalf("wrapped large response reader type mismatch: %T", finalizingReader.ReadCloser) + } + + body, err := io.ReadAll(reader) + if err != nil { + t.Fatalf("failed to read passthrough body: %v", err) + } + if !strings.Contains(string(body), `"message_id":"resp-large"`) { + t.Fatalf("passthrough body mismatch: %s", body) + } + if got := finalizerCalls.Load(); got != 0 { + t.Fatalf("finalizer ran before passthrough reader close: got %d calls", got) + } + if err := reader.Close(); err != nil { + t.Fatalf("failed to close large response reader: %v", err) + } + if got := finalizerCalls.Load(); got != 1 { + t.Fatalf("finalizer calls mismatch after passthrough delivery: got %d, want 1", got) + } + if ended, _ := ctx.Value(schemas.BifrostContextKeyStreamEndIndicator).(bool); !ended { + t.Fatal("passthrough stream was not marked complete before finalization") + } + if largeReader.Resp != nil { + t.Fatal("large response reader did not release its fasthttp response") + } +} + +func testGigaChatResponsesRequest() *schemas.BifrostResponsesRequest { + return &schemas.BifrostResponsesRequest{ + Model: "GigaChat-2", + Input: []schemas.ResponsesMessage{{ + Role: schemas.Ptr(schemas.ResponsesInputMessageRoleUser), + Content: &schemas.ResponsesMessageContent{ContentStr: schemas.Ptr("hello")}, + }}, + } +} + +func testGigaChatResponsesExecutionRequest() *schemas.BifrostResponsesRequest { + maxTokens := 128 + return &schemas.BifrostResponsesRequest{ + Model: "GigaChat-2", + Input: []schemas.ResponsesMessage{{ + Role: schemas.Ptr(schemas.ResponsesInputMessageRoleUser), + Content: &schemas.ResponsesMessageContent{ContentStr: schemas.Ptr("Привет")}, + }}, + Params: &schemas.ResponsesParameters{ + MaxOutputTokens: &maxTokens, + }, + } +} + +func assertGigaChatResponsesRequestBody(t *testing.T, request *http.Request) { + t.Helper() + + if request.Method != http.MethodPost { + t.Fatalf("method mismatch: got %s, want POST", request.Method) + } + if got := request.Header.Get("Content-Type"); !strings.HasPrefix(got, "application/json") { + t.Fatalf("content type mismatch: got %q", got) + } + if got := request.Header.Get("Accept"); got != "application/json" { + t.Fatalf("accept header mismatch: got %q", got) + } + if got := request.Header.Get(gigaChatUserAgentHeader); got != gigaChatUserAgent { + t.Fatalf("user-agent mismatch: got %q, want %q", got, gigaChatUserAgent) + } + + body, err := io.ReadAll(request.Body) + if err != nil { + t.Fatalf("failed to read request body: %v", err) + } + var payload map[string]interface{} + if err := json.Unmarshal(body, &payload); err != nil { + t.Fatalf("failed to unmarshal request body %s: %v", body, err) + } + if got := payload["model"]; got != "GigaChat-2" { + t.Fatalf("model mismatch: got %#v", got) + } + if _, ok := payload["stream"]; ok { + t.Fatalf("non-streaming responses request should omit stream: %s", body) + } + modelOptions, ok := payload["model_options"].(map[string]interface{}) + if !ok { + t.Fatalf("model_options mismatch: %#v", payload["model_options"]) + } + if got := modelOptions["max_tokens"]; got != float64(128) { + t.Fatalf("max_tokens mismatch: got %#v", got) + } + storage, ok := payload["storage"].(map[string]interface{}) + if !ok || len(storage) != 0 { + t.Fatalf("storage mismatch: %#v", payload["storage"]) + } + messages, ok := payload["messages"].([]interface{}) + if !ok || len(messages) != 1 { + t.Fatalf("messages mismatch: %#v", payload["messages"]) + } + message, ok := messages[0].(map[string]interface{}) + if !ok { + t.Fatalf("message shape mismatch: %#v", messages[0]) + } + if got := message["role"]; got != "user" { + t.Fatalf("message role mismatch: got %#v", got) + } + content, ok := message["content"].([]interface{}) + if !ok || len(content) != 1 { + t.Fatalf("message content mismatch: %#v", message["content"]) + } + contentPart, ok := content[0].(map[string]interface{}) + if !ok || contentPart["text"] != "Привет" { + t.Fatalf("content part mismatch: %#v", content[0]) + } +} + +func assertGigaChatResponsesBodyFile(t *testing.T, body []byte, wantFileID string) map[string]interface{} { + t.Helper() + + var payload map[string]interface{} + if err := json.Unmarshal(body, &payload); err != nil { + t.Fatalf("failed to unmarshal responses body %s: %v", body, err) + } + messages, ok := payload["messages"].([]interface{}) + if !ok || len(messages) != 1 { + t.Fatalf("messages mismatch: %#v", payload["messages"]) + } + message, ok := messages[0].(map[string]interface{}) + if !ok { + t.Fatalf("message shape mismatch: %#v", messages[0]) + } + content, ok := message["content"].([]interface{}) + if !ok || len(content) != 2 { + t.Fatalf("message content mismatch: %#v", message["content"]) + } + filePart, ok := content[1].(map[string]interface{}) + if !ok { + t.Fatalf("file content shape mismatch: %#v", content[1]) + } + files, ok := filePart["files"].([]interface{}) + if !ok || len(files) != 1 { + t.Fatalf("files mismatch: %#v", filePart["files"]) + } + file, ok := files[0].(map[string]interface{}) + if !ok { + t.Fatalf("file shape mismatch: %#v", files[0]) + } + if got := file["id"]; got != wantFileID { + t.Fatalf("file id mismatch: got %#v, want %q body %s", got, wantFileID, body) + } + return payload +} + +func assertGigaChatResponsesStreamRequestBody(t *testing.T, request *http.Request) { + t.Helper() + + if request.Method != http.MethodPost { + t.Fatalf("method mismatch: got %s, want POST", request.Method) + } + if got := request.Header.Get("Content-Type"); !strings.HasPrefix(got, "application/json") { + t.Fatalf("content type mismatch: got %q", got) + } + if got := request.Header.Get("Accept"); got != "text/event-stream" { + t.Fatalf("accept header mismatch: got %q", got) + } + if got := request.Header.Get(gigaChatUserAgentHeader); got != gigaChatUserAgent { + t.Fatalf("user-agent mismatch: got %q, want %q", got, gigaChatUserAgent) + } + + body, err := io.ReadAll(request.Body) + if err != nil { + t.Fatalf("failed to read request body: %v", err) + } + var payload map[string]interface{} + if err := json.Unmarshal(body, &payload); err != nil { + t.Fatalf("failed to unmarshal request body %s: %v", body, err) + } + if got := payload["model"]; got != "GigaChat-2" { + t.Fatalf("model mismatch: got %#v", got) + } + if got := payload["stream"]; got != true { + t.Fatalf("stream mismatch: got %#v, want true; body=%s", got, body) + } + messages, ok := payload["messages"].([]interface{}) + if !ok || len(messages) != 1 { + t.Fatalf("messages mismatch: %#v", payload["messages"]) + } +} + +func collectGigaChatResponsesStreamResponses(t *testing.T, chunks []*schemas.BifrostStreamChunk) []*schemas.BifrostResponsesStreamResponse { + t.Helper() + + responses := make([]*schemas.BifrostResponsesStreamResponse, 0, len(chunks)) + for _, chunk := range chunks { + if chunk == nil { + t.Fatal("got nil stream chunk") + } + if chunk.BifrostError != nil { + t.Fatalf("unexpected stream error: %v", chunk.BifrostError) + } + if chunk.BifrostResponsesStreamResponse == nil { + t.Fatalf("missing responses stream response: %#v", chunk) + } + if chunk.BifrostResponsesStreamResponse.ExtraFields.Provider != schemas.GigaChat { + t.Fatalf("provider mismatch: got %q, want %q", chunk.BifrostResponsesStreamResponse.ExtraFields.Provider, schemas.GigaChat) + } + if chunk.BifrostResponsesStreamResponse.Response != nil && + chunk.BifrostResponsesStreamResponse.Response.ExtraFields.Provider != schemas.GigaChat { + t.Fatalf("nested response provider mismatch: got %q, want %q", chunk.BifrostResponsesStreamResponse.Response.ExtraFields.Provider, schemas.GigaChat) + } + responses = append(responses, chunk.BifrostResponsesStreamResponse) + } + return responses +} + +func assertGigaChatResponsesStreamTypes(t *testing.T, responses []*schemas.BifrostResponsesStreamResponse, want []schemas.ResponsesStreamResponseType) { + t.Helper() + + if len(responses) != len(want) { + gotTypes := make([]schemas.ResponsesStreamResponseType, 0, len(responses)) + for _, response := range responses { + gotTypes = append(gotTypes, response.Type) + } + t.Fatalf("response type count mismatch: got %d %v, want %d %v", len(responses), gotTypes, len(want), want) + } + for index, wantType := range want { + if responses[index].Type != wantType { + t.Fatalf("response type[%d] mismatch: got %q, want %q", index, responses[index].Type, wantType) + } + } +} + +func mustGigaChatToolParameters(t *testing.T, raw string) *schemas.ToolFunctionParameters { + t.Helper() + + var parameters schemas.ToolFunctionParameters + if err := json.Unmarshal([]byte(raw), ¶meters); err != nil { + t.Fatalf("failed to unmarshal tool parameters: %v", err) + } + return ¶meters +} diff --git a/core/providers/gigachat/tools_test.go b/core/providers/gigachat/tools_test.go new file mode 100644 index 00000000000..e22efbf371c --- /dev/null +++ b/core/providers/gigachat/tools_test.go @@ -0,0 +1,1104 @@ +package gigachat + +import ( + "encoding/json" + "strings" + "testing" + + "github.com/maximhq/bifrost/core/schemas" +) + +func TestGigaChatTools(t *testing.T) { + testGigaChatTools(t) +} + +func testGigaChatTools(t *testing.T) { + t.Parallel() + + t.Run("ChatMapsFunctionToolsAndHistory", testGigaChatToolsChatMapsFunctionToolsAndHistory) + t.Run("ChatToolChoiceVariants", testGigaChatToolsChatToolChoiceVariants) + t.Run("ChatDeduplicatesEquivalentFunctionTools", testGigaChatToolsChatDeduplicatesEquivalentFunctionTools) + t.Run("ChatRejectsUnsupportedPolicy", testGigaChatToolsChatRejectsUnsupportedPolicy) + t.Run("ResponsesMapsBuiltInTools", testGigaChatToolsResponsesMapsBuiltInTools) + t.Run("ResponsesOmitNestedBuiltInToolTypes", testGigaChatToolsResponsesOmitNestedBuiltInToolTypes) + t.Run("ResponsesToolChoiceVariants", testGigaChatToolsResponsesToolChoiceVariants) + t.Run("ResponsesRemapsReservedFunctionNames", testGigaChatToolsResponsesRemapsReservedFunctionNames) + t.Run("ResponsesDeduplicatesEquivalentFunctionTools", testGigaChatToolsResponsesDeduplicatesEquivalentFunctionTools) + t.Run("SanitizesFunctionSchemas", testGigaChatToolsSanitizesFunctionSchemas) + t.Run("ResponsesRejectsUnsupportedPolicy", testGigaChatToolsResponsesRejectsUnsupportedPolicy) +} + +func testGigaChatToolsChatMapsFunctionToolsAndHistory(t *testing.T) { + t.Parallel() + + toolName := "get_weather" + toolCallID := "state-weather" + toolCallType := string(schemas.ChatToolTypeFunction) + toolArguments := `{"city":"Moscow"}` + reasoning := "I should call get_weather" + result := `{"temperature":5}` + request := &schemas.BifrostChatRequest{ + Model: "GigaChat", + Input: []schemas.ChatMessage{ + { + Role: schemas.ChatMessageRoleUser, + Content: &schemas.ChatMessageContent{ContentStr: schemas.Ptr("Weather?")}, + }, + { + Role: schemas.ChatMessageRoleAssistant, + Content: &schemas.ChatMessageContent{ContentStr: schemas.Ptr("")}, + ChatAssistantMessage: &schemas.ChatAssistantMessage{ + Reasoning: &reasoning, + ToolCalls: []schemas.ChatAssistantMessageToolCall{{ + Type: &toolCallType, + ID: &toolCallID, + Function: schemas.ChatAssistantMessageToolCallFunction{ + Name: &toolName, + Arguments: toolArguments, + }, + }}, + }, + }, + { + Role: schemas.ChatMessageRoleTool, + Content: &schemas.ChatMessageContent{ContentStr: &result}, + ChatToolMessage: &schemas.ChatToolMessage{ToolCallID: &toolCallID}, + }, + { + Role: schemas.ChatMessageRoleUser, + Content: &schemas.ChatMessageContent{ContentStr: schemas.Ptr("What should I wear?")}, + }, + }, + Params: &schemas.ChatParameters{ + Tools: []schemas.ChatTool{testGigaChatChatFunctionTool(t, toolName)}, + ToolChoice: &schemas.ChatToolChoice{ + ChatToolChoiceStruct: &schemas.ChatToolChoiceStruct{ + Type: schemas.ChatToolChoiceTypeFunction, + Function: &schemas.ChatToolChoiceFunction{Name: toolName}, + }, + }, + }, + } + + gigaChatReq, err := ToGigaChatChatRequest(testBifrostContext(), request) + if err != nil { + t.Fatalf("ToGigaChatChatRequest returned error: %v", err) + } + if len(gigaChatReq.Functions) != 1 || gigaChatReq.Functions[0].Name != toolName { + t.Fatalf("functions mismatch: %#v", gigaChatReq.Functions) + } + choice, ok := gigaChatReq.FunctionCall.(GigaChatFunctionCallChoice) + if !ok || choice.Name != toolName { + t.Fatalf("function_call mismatch: %#v", gigaChatReq.FunctionCall) + } + if len(gigaChatReq.Messages) != 4 { + t.Fatalf("message count mismatch: got %d", len(gigaChatReq.Messages)) + } + assistant := gigaChatReq.Messages[1] + if assistant.FunctionCall == nil || assistant.FunctionCall.Name != toolName || string(assistant.FunctionCall.Arguments) != toolArguments { + t.Fatalf("assistant function_call mismatch: %#v", assistant) + } + if assistant.Reasoning == nil || *assistant.Reasoning != reasoning { + t.Fatalf("assistant reasoning_content mismatch: %#v", assistant.Reasoning) + } + if assistant.FunctionsStateID == nil || *assistant.FunctionsStateID != toolCallID { + t.Fatalf("functions_state_id mismatch: %#v", assistant.FunctionsStateID) + } + functionResult := gigaChatReq.Messages[2] + if functionResult.Role != "function" || functionResult.Name == nil || *functionResult.Name != toolName { + t.Fatalf("function result message mismatch: %#v", functionResult) + } + if functionResult.Content == nil || functionResult.Content.ContentStr == nil || *functionResult.Content.ContentStr != result { + t.Fatalf("function result content mismatch: %#v", functionResult.Content) + } +} + +func testGigaChatToolsChatToolChoiceVariants(t *testing.T) { + t.Parallel() + + toolName := "get_weather" + tests := []struct { + name string + choice *schemas.ChatToolChoice + wantMode string + wantForced string + }{ + { + name: "StringAuto", + choice: &schemas.ChatToolChoice{ChatToolChoiceStr: schemas.Ptr("auto")}, + wantMode: "auto", + }, + { + name: "StringRequired", + choice: &schemas.ChatToolChoice{ChatToolChoiceStr: schemas.Ptr("required")}, + wantForced: toolName, + }, + { + name: "StringAny", + choice: &schemas.ChatToolChoice{ChatToolChoiceStr: schemas.Ptr("any")}, + wantForced: toolName, + }, + { + name: "StringNone", + choice: &schemas.ChatToolChoice{ChatToolChoiceStr: schemas.Ptr("none")}, + wantMode: "none", + }, + { + name: "StructAuto", + choice: &schemas.ChatToolChoice{ChatToolChoiceStruct: &schemas.ChatToolChoiceStruct{ + Type: schemas.ChatToolChoiceTypeAuto, + }}, + wantMode: "auto", + }, + { + name: "StructNone", + choice: &schemas.ChatToolChoice{ChatToolChoiceStruct: &schemas.ChatToolChoiceStruct{ + Type: schemas.ChatToolChoiceTypeNone, + }}, + wantMode: "none", + }, + { + name: "StructRequired", + choice: &schemas.ChatToolChoice{ChatToolChoiceStruct: &schemas.ChatToolChoiceStruct{ + Type: schemas.ChatToolChoiceTypeRequired, + }}, + wantForced: toolName, + }, + { + name: "StructAny", + choice: &schemas.ChatToolChoice{ChatToolChoiceStruct: &schemas.ChatToolChoiceStruct{ + Type: schemas.ChatToolChoiceTypeAny, + }}, + wantForced: toolName, + }, + { + name: "StructFunction", + choice: &schemas.ChatToolChoice{ChatToolChoiceStruct: &schemas.ChatToolChoiceStruct{ + Type: schemas.ChatToolChoiceTypeFunction, + Function: &schemas.ChatToolChoiceFunction{Name: toolName}, + }}, + wantForced: toolName, + }, + } + + for _, test := range tests { + test := test + t.Run(test.name, func(t *testing.T) { + t.Parallel() + + request := testGigaChatChatToolRequest(t, toolName) + request.Params.ToolChoice = test.choice + gigaChatReq, err := ToGigaChatChatRequest(testBifrostContext(), request) + if err != nil { + t.Fatalf("ToGigaChatChatRequest returned error: %v", err) + } + if test.wantForced != "" { + choice, ok := gigaChatReq.FunctionCall.(GigaChatFunctionCallChoice) + if !ok || choice.Name != test.wantForced { + t.Fatalf("forced function_call mismatch: %#v", gigaChatReq.FunctionCall) + } + return + } + mode, ok := gigaChatReq.FunctionCall.(string) + if !ok || mode != test.wantMode { + t.Fatalf("function_call mode mismatch: got %#v, want %q", gigaChatReq.FunctionCall, test.wantMode) + } + }) + } +} + +func testGigaChatToolsChatDeduplicatesEquivalentFunctionTools(t *testing.T) { + t.Parallel() + + request := testGigaChatChatToolRequest(t, "get_weather") + request.Params.Tools = append(request.Params.Tools, request.Params.Tools[0]) + + gigaChatReq, err := ToGigaChatChatRequest(testBifrostContext(), request) + if err != nil { + t.Fatalf("ToGigaChatChatRequest returned error: %v", err) + } + if len(gigaChatReq.Functions) != 1 || gigaChatReq.Functions[0].Name != "get_weather" { + t.Fatalf("duplicate function tools were not deduplicated: %#v", gigaChatReq.Functions) + } +} + +func testGigaChatToolsChatRejectsUnsupportedPolicy(t *testing.T) { + t.Parallel() + + parallelToolCalls := true + strict := true + timeTool := testGigaChatChatFunctionTool(t, "get_time") + tests := []struct { + name string + mutate func(*schemas.BifrostChatRequest) + wantErr string + }{ + { + name: "InvalidJSONSchema", + mutate: func(request *schemas.BifrostChatRequest) { + request.Params.Tools[0].Function.Parameters = invalidGigaChatToolParameters() + }, + wantErr: "JSON schema is invalid", + }, + { + name: "GigaChatBuiltInName", + mutate: func(request *schemas.BifrostChatRequest) { + request.Params.Tools[0].Function.Name = "text2image" + }, + wantErr: "built-in function", + }, + { + name: "OpenAIStrictMode", + mutate: func(request *schemas.BifrostChatRequest) { + request.Params.Tools[0].Function.Strict = &strict + }, + wantErr: "strict mode", + }, + { + name: "CustomTool", + mutate: func(request *schemas.BifrostChatRequest) { + request.Params.Tools = []schemas.ChatTool{{ + Type: schemas.ChatToolTypeCustom, + Name: "custom_tool", + Custom: &schemas.ChatToolCustom{}, + }} + }, + wantErr: "function tools only", + }, + { + name: "ParallelToolCalls", + mutate: func(request *schemas.BifrostChatRequest) { + request.Params.ParallelToolCalls = ¶llelToolCalls + }, + wantErr: "parallel_tool_calls", + }, + { + name: "RequiredToolChoiceWithMultipleFunctions", + mutate: func(request *schemas.BifrostChatRequest) { + request.Params.Tools = append(request.Params.Tools, timeTool) + request.Params.ToolChoice = &schemas.ChatToolChoice{ChatToolChoiceStr: schemas.Ptr("required")} + }, + wantErr: "cannot require an arbitrary", + }, + { + name: "AutoToolChoiceWithoutFunctions", + mutate: func(request *schemas.BifrostChatRequest) { + request.Params.Tools = nil + request.Params.ToolChoice = &schemas.ChatToolChoice{ChatToolChoiceStr: schemas.Ptr("auto")} + }, + wantErr: "requires at least one", + }, + { + name: "ConflictingDuplicateFunction", + mutate: func(request *schemas.BifrostChatRequest) { + duplicate := request.Params.Tools[0] + function := *duplicate.Function + duplicate.Function = &function + duplicate.Function.Description = schemas.Ptr("Different weather function.") + request.Params.Tools = append(request.Params.Tools, duplicate) + }, + wantErr: "different definition", + }, + { + name: "ExtraParamFunctionsBypass", + mutate: func(request *schemas.BifrostChatRequest) { + request.Params.ExtraParams = map[string]interface{}{"functions": []interface{}{}} + }, + wantErr: "extra_params.functions", + }, + } + + for _, test := range tests { + test := test + t.Run(test.name, func(t *testing.T) { + t.Parallel() + + request := testGigaChatChatToolRequest(t, "get_weather") + test.mutate(request) + _, err := ToGigaChatChatRequest(testBifrostContext(), request) + if err == nil || !strings.Contains(err.Error(), test.wantErr) { + t.Fatalf("expected error containing %q, got %v", test.wantErr, err) + } + }) + } +} + +func testGigaChatToolsResponsesMapsBuiltInTools(t *testing.T) { + t.Parallel() + + functionName := "web_search" + imageModel := "Kandinsky" + imageSize := "1024x1024" + searchContextSize := "high" + city := "Moscow" + country := "RU" + maxContentTokens := 2000 + request := testGigaChatResponsesToolRequest(t, functionName) + request.Params.Tools = append(request.Params.Tools, + schemas.ResponsesTool{ + Type: schemas.ResponsesToolTypeWebSearchPreview, + ResponsesToolWebSearchPreview: &schemas.ResponsesToolWebSearchPreview{ + SearchContextSize: &searchContextSize, + UserLocation: &schemas.ResponsesToolWebSearchUserLocation{ + City: &city, + Country: &country, + }, + }, + }, + schemas.ResponsesTool{ + Type: schemas.ResponsesToolTypeCodeInterpreter, + ResponsesToolCodeInterpreter: &schemas.ResponsesToolCodeInterpreter{ + Container: map[string]interface{}{"type": "auto"}, + }, + }, + schemas.ResponsesTool{ + Type: schemas.ResponsesToolTypeImageGeneration, + ResponsesToolImageGeneration: &schemas.ResponsesToolImageGeneration{ + Model: &imageModel, + Size: &imageSize, + }, + }, + schemas.ResponsesTool{ + Type: schemas.ResponsesToolTypeWebFetch, + ResponsesToolWebFetch: &schemas.ResponsesToolWebFetch{ + MaxContentTokens: &maxContentTokens, + }, + }, + schemas.ResponsesTool{ + Type: schemas.ResponsesToolType("model_3d_generate"), + Name: schemas.Ptr("make_model"), + Description: schemas.Ptr("Generate a 3D model."), + }, + ) + + gigaChatReq, err := ToGigaChatResponsesRequest(request) + if err != nil { + t.Fatalf("ToGigaChatResponsesRequest returned error: %v", err) + } + if len(gigaChatReq.Tools) != 6 { + t.Fatalf("tool count mismatch: got %d, tools=%#v", len(gigaChatReq.Tools), gigaChatReq.Tools) + } + functionsTool := gigaChatReq.Tools[0] + if functionsTool.Functions == nil || len(functionsTool.Functions.Specifications) != 1 { + t.Fatalf("functions tool mismatch: %#v", functionsTool) + } + if got := functionsTool.Functions.Specifications[0].Name; got != "__bifrost_gigachat_user_web_search" { + t.Fatalf("function remap mismatch: got %q", got) + } + + webSearchTool := gigaChatReq.Tools[1] + if webSearchTool.WebSearch == nil { + t.Fatalf("web_search tool mismatch: %#v", webSearchTool) + } + if webSearchTool.WebSearch.Type != nil { + t.Fatalf("web_search should not include nested type, got %#v", webSearchTool.WebSearch.Type) + } + if len(webSearchTool.WebSearch.Flags) != 1 || webSearchTool.WebSearch.Flags[0] != "search_context_size:high" { + t.Fatalf("web_search flags mismatch: %#v", webSearchTool.WebSearch.Flags) + } + userLocation, ok := gigaChatReq.UserInfo["user_location"].(map[string]interface{}) + if !ok || userLocation["city"] != city || userLocation["country"] != country { + t.Fatalf("user_info mismatch: %#v", gigaChatReq.UserInfo) + } + + codeTool := gigaChatReq.Tools[2] + container, ok := codeTool.CodeInterpreter["container"].(map[string]interface{}) + if !ok || container["type"] != "auto" { + t.Fatalf("code_interpreter container mismatch: %#v", codeTool.CodeInterpreter) + } + + imageTool := gigaChatReq.Tools[3] + if imageTool.ImageGenerate["type"] != nil || + imageTool.ImageGenerate["model"] != imageModel || + imageTool.ImageGenerate["size"] != imageSize { + t.Fatalf("image_generate config mismatch: %#v", imageTool.ImageGenerate) + } + + urlTool := gigaChatReq.Tools[4] + if urlTool.URLContentExtraction["type"] != nil || + urlTool.URLContentExtraction["max_content_tokens"] != float64(maxContentTokens) { + t.Fatalf("url_content_extraction config mismatch: %#v", urlTool.URLContentExtraction) + } + + model3DTool := gigaChatReq.Tools[5] + if model3DTool.Model3DGenerate["type"] != nil || + model3DTool.Model3DGenerate["name"] != "make_model" || + model3DTool.Model3DGenerate["description"] != "Generate a 3D model." { + t.Fatalf("model_3d_generate config mismatch: %#v", model3DTool.Model3DGenerate) + } +} + +func testGigaChatToolsResponsesOmitNestedBuiltInToolTypes(t *testing.T) { + t.Parallel() + + request := &schemas.BifrostResponsesRequest{ + Model: "GigaChat-2", + Input: []schemas.ResponsesMessage{{ + Role: schemas.Ptr(schemas.ResponsesInputMessageRoleUser), + Content: &schemas.ResponsesMessageContent{ContentStr: schemas.Ptr("Draw space.")}, + }}, + Params: &schemas.ResponsesParameters{ + Tools: []schemas.ResponsesTool{ + {Type: schemas.ResponsesToolTypeImageGeneration}, + {Type: schemas.ResponsesToolTypeWebSearch}, + {Type: schemas.ResponsesToolTypeCodeInterpreter}, + {Type: schemas.ResponsesToolTypeWebFetch}, + {Type: schemas.ResponsesToolType("model_3d_generate")}, + }, + }, + } + + gigaChatReq, err := ToGigaChatResponsesRequest(request) + if err != nil { + t.Fatalf("ToGigaChatResponsesRequest returned error: %v", err) + } + raw, err := json.Marshal(gigaChatReq) + if err != nil { + t.Fatalf("failed to marshal GigaChat request: %v", err) + } + body := string(raw) + wantTools := `"tools":[{"image_generate":{}},{"web_search":{}},{"code_interpreter":{}},{"url_content_extraction":{}},{"model_3d_generate":{}}]` + if !strings.Contains(body, wantTools) { + t.Fatalf("built-in tool JSON mismatch:\n got: %s\nwant substring: %s", body, wantTools) + } + if strings.Contains(body, `"image_generate":{"type":`) || + strings.Contains(body, `"web_search":{"type":`) || + strings.Contains(body, `"code_interpreter":{"type":`) || + strings.Contains(body, `"url_content_extraction":{"type":`) || + strings.Contains(body, `"model_3d_generate":{"type":`) { + t.Fatalf("built-in tool JSON should not include nested type discriminators: %s", body) + } +} + +func testGigaChatToolsResponsesToolChoiceVariants(t *testing.T) { + t.Parallel() + + toolName := "get_weather" + tests := []struct { + name string + choice *schemas.ResponsesToolChoice + mutate func(*schemas.BifrostResponsesRequest) + wantMode string + wantFunction string + wantTool string + }{ + { + name: "StringAuto", + choice: &schemas.ResponsesToolChoice{ResponsesToolChoiceStr: schemas.Ptr("auto")}, + wantMode: "auto", + }, + { + name: "StringRequired", + choice: &schemas.ResponsesToolChoice{ResponsesToolChoiceStr: schemas.Ptr("required")}, + wantFunction: toolName, + }, + { + name: "StringAny", + choice: &schemas.ResponsesToolChoice{ResponsesToolChoiceStr: schemas.Ptr("any")}, + wantFunction: toolName, + }, + { + name: "StringNone", + choice: &schemas.ResponsesToolChoice{ResponsesToolChoiceStr: schemas.Ptr("none")}, + wantMode: "none", + }, + { + name: "StructAuto", + choice: &schemas.ResponsesToolChoice{ResponsesToolChoiceStruct: &schemas.ResponsesToolChoiceStruct{ + Type: schemas.ResponsesToolChoiceTypeAuto, + }}, + wantMode: "auto", + }, + { + name: "StructNone", + choice: &schemas.ResponsesToolChoice{ResponsesToolChoiceStruct: &schemas.ResponsesToolChoiceStruct{ + Type: schemas.ResponsesToolChoiceTypeNone, + }}, + wantMode: "none", + }, + { + name: "StructRequired", + choice: &schemas.ResponsesToolChoice{ResponsesToolChoiceStruct: &schemas.ResponsesToolChoiceStruct{ + Type: schemas.ResponsesToolChoiceTypeRequired, + }}, + wantFunction: toolName, + }, + { + name: "StructAny", + choice: &schemas.ResponsesToolChoice{ResponsesToolChoiceStruct: &schemas.ResponsesToolChoiceStruct{ + Type: schemas.ResponsesToolChoiceTypeAny, + }}, + wantFunction: toolName, + }, + { + name: "StructFunction", + choice: &schemas.ResponsesToolChoice{ResponsesToolChoiceStruct: &schemas.ResponsesToolChoiceStruct{ + Type: schemas.ResponsesToolChoiceTypeFunction, + Name: &toolName, + }}, + wantFunction: toolName, + }, + { + name: "StructReservedFunction", + choice: &schemas.ResponsesToolChoice{ResponsesToolChoiceStruct: &schemas.ResponsesToolChoiceStruct{ + Type: schemas.ResponsesToolChoiceTypeFunction, + Name: schemas.Ptr("web_search"), + }}, + mutate: func(request *schemas.BifrostResponsesRequest) { + request.Params.Tools[0].Name = schemas.Ptr("web_search") + }, + wantFunction: "__bifrost_gigachat_user_web_search", + }, + { + name: "StringAutoWithBuiltInOnly", + choice: &schemas.ResponsesToolChoice{ + ResponsesToolChoiceStr: schemas.Ptr("auto"), + }, + mutate: func(request *schemas.BifrostResponsesRequest) { + request.Params.Tools = []schemas.ResponsesTool{{ + Type: schemas.ResponsesToolTypeCodeInterpreter, + ResponsesToolCodeInterpreter: &schemas.ResponsesToolCodeInterpreter{Container: map[string]interface{}{"type": "auto"}}, + }} + }, + wantMode: "auto", + }, + { + name: "StructCodeInterpreter", + choice: &schemas.ResponsesToolChoice{ResponsesToolChoiceStruct: &schemas.ResponsesToolChoiceStruct{ + Type: schemas.ResponsesToolChoiceTypeCodeInterpreter, + }}, + mutate: func(request *schemas.BifrostResponsesRequest) { + request.Params.Tools = []schemas.ResponsesTool{{ + Type: schemas.ResponsesToolTypeCodeInterpreter, + ResponsesToolCodeInterpreter: &schemas.ResponsesToolCodeInterpreter{Container: map[string]interface{}{"type": "auto"}}, + }} + }, + wantTool: "code_interpreter", + }, + { + name: "StructImageGeneration", + choice: &schemas.ResponsesToolChoice{ResponsesToolChoiceStruct: &schemas.ResponsesToolChoiceStruct{ + Type: schemas.ResponsesToolChoiceTypeImageGeneration, + }}, + mutate: func(request *schemas.BifrostResponsesRequest) { + request.Params.Tools = []schemas.ResponsesTool{{ + Type: schemas.ResponsesToolTypeImageGeneration, + ResponsesToolImageGeneration: &schemas.ResponsesToolImageGeneration{ + Size: schemas.Ptr("1024x1024"), + }, + }} + }, + wantTool: "image_generate", + }, + { + name: "StructWebSearchPreview", + choice: &schemas.ResponsesToolChoice{ResponsesToolChoiceStruct: &schemas.ResponsesToolChoiceStruct{ + Type: schemas.ResponsesToolChoiceTypeWebSearchPreview, + }}, + mutate: func(request *schemas.BifrostResponsesRequest) { + request.Params.Tools = []schemas.ResponsesTool{{ + Type: schemas.ResponsesToolTypeWebSearchPreview, + ResponsesToolWebSearchPreview: &schemas.ResponsesToolWebSearchPreview{ + SearchContextSize: schemas.Ptr("low"), + }, + }} + }, + wantTool: "web_search", + }, + { + name: "StructURLContentExtraction", + choice: &schemas.ResponsesToolChoice{ResponsesToolChoiceStruct: &schemas.ResponsesToolChoiceStruct{ + Type: schemas.ResponsesToolChoiceType("url_content_extraction"), + }}, + mutate: func(request *schemas.BifrostResponsesRequest) { + request.Params.Tools = []schemas.ResponsesTool{{ + Type: schemas.ResponsesToolType("url_content_extraction"), + }} + }, + wantTool: "url_content_extraction", + }, + { + name: "StructModel3DGenerate", + choice: &schemas.ResponsesToolChoice{ResponsesToolChoiceStruct: &schemas.ResponsesToolChoiceStruct{ + Type: schemas.ResponsesToolChoiceType("model_3d_generate"), + }}, + mutate: func(request *schemas.BifrostResponsesRequest) { + request.Params.Tools = []schemas.ResponsesTool{{ + Type: schemas.ResponsesToolType("model_3d_generate"), + }} + }, + wantTool: "model_3d_generate", + }, + } + + for _, test := range tests { + test := test + t.Run(test.name, func(t *testing.T) { + t.Parallel() + + request := testGigaChatResponsesToolRequest(t, toolName) + if test.mutate != nil { + test.mutate(request) + } + request.Params.ToolChoice = test.choice + gigaChatReq, err := ToGigaChatResponsesRequest(request) + if err != nil { + t.Fatalf("ToGigaChatResponsesRequest returned error: %v", err) + } + if test.wantFunction != "" { + if gigaChatReq.ToolConfig == nil || gigaChatReq.ToolConfig.FunctionName == nil || *gigaChatReq.ToolConfig.FunctionName != test.wantFunction || gigaChatReq.ToolConfig.Mode != "forced" { + t.Fatalf("forced tool_config mismatch: %#v", gigaChatReq.ToolConfig) + } + return + } + if test.wantTool != "" { + if gigaChatReq.ToolConfig == nil || gigaChatReq.ToolConfig.ToolName == nil || *gigaChatReq.ToolConfig.ToolName != test.wantTool || gigaChatReq.ToolConfig.Mode != "forced" { + t.Fatalf("forced built-in tool_config mismatch: %#v", gigaChatReq.ToolConfig) + } + return + } + if gigaChatReq.ToolConfig == nil || gigaChatReq.ToolConfig.Mode != test.wantMode { + t.Fatalf("tool_config mode mismatch: got %#v, want %q", gigaChatReq.ToolConfig, test.wantMode) + } + }) + } +} + +func testGigaChatToolsResponsesRemapsReservedFunctionNames(t *testing.T) { + t.Parallel() + + functionName := "web_search" + arguments := `{"query":"GigaChat"}` + request := testGigaChatResponsesToolRequest(t, functionName) + request.Input = append(request.Input, schemas.ResponsesMessage{ + Type: schemas.Ptr(schemas.ResponsesMessageTypeFunctionCall), + ResponsesToolMessage: &schemas.ResponsesToolMessage{ + Name: &functionName, + CallID: schemas.Ptr("state-web-search"), + Arguments: &arguments, + }, + }) + + gigaChatReq, err := ToGigaChatResponsesRequest(request) + if err != nil { + t.Fatalf("ToGigaChatResponsesRequest returned error: %v", err) + } + if got := gigaChatReq.Tools[0].Functions.Specifications[0].Name; got != "__bifrost_gigachat_user_web_search" { + t.Fatalf("function specification name mismatch: got %q", got) + } + if gigaChatReq.Messages[1].FunctionCall == nil || gigaChatReq.Messages[1].FunctionCall.Name != "__bifrost_gigachat_user_web_search" { + t.Fatalf("function call remap mismatch: %#v", gigaChatReq.Messages[1].FunctionCall) + } + + response := ToBifrostResponsesResponse(schemas.GigaChat, &GigaChatResponsesResponse{ + ID: "resp-1", + Model: "GigaChat-2", + Messages: []GigaChatResponsesMessage{{ + Role: "assistant", + ToolsStateID: schemas.Ptr("state-web-search"), + Content: []GigaChatResponsesContentPart{{ + FunctionCall: &GigaChatResponsesFunctionCall{ + Name: "__bifrost_gigachat_user_web_search", + Arguments: map[string]interface{}{"query": "GigaChat"}, + }, + }}, + }}, + }) + if response == nil || len(response.Output) != 1 || response.Output[0].ResponsesToolMessage == nil || response.Output[0].ResponsesToolMessage.Name == nil { + t.Fatalf("response output mismatch: %#v", response) + } + if got := *response.Output[0].ResponsesToolMessage.Name; got != functionName { + t.Fatalf("function response name mismatch: got %q", got) + } +} + +func testGigaChatToolsResponsesDeduplicatesEquivalentFunctionTools(t *testing.T) { + t.Parallel() + + request := testGigaChatResponsesToolRequest(t, "get_horoscope") + for range 19 { + request.Params.Tools = append(request.Params.Tools, request.Params.Tools[0]) + } + + gigaChatReq, err := ToGigaChatResponsesRequest(request) + if err != nil { + t.Fatalf("ToGigaChatResponsesRequest returned error: %v", err) + } + if len(gigaChatReq.Tools) != 1 || gigaChatReq.Tools[0].Functions == nil || len(gigaChatReq.Tools[0].Functions.Specifications) != 1 { + t.Fatalf("duplicate function tools were not deduplicated: %#v", gigaChatReq.Tools) + } + if got := gigaChatReq.Tools[0].Functions.Specifications[0].Name; got != "get_horoscope" { + t.Fatalf("function specification name mismatch: got %q", got) + } +} + +func testGigaChatToolsSanitizesFunctionSchemas(t *testing.T) { + t.Parallel() + + rawSchema := `{ + "type": "object", + "$defs": { + "Location": { + "type": "object", + "properties": { + "city": {"anyOf": [{"type": "string"}, {"type": "null"}]}, + "coords": {"type": ["object", "null"]} + } + } + }, + "properties": { + "nickname": {"type": ["string", "null"], "nullable": true}, + "location": {"$ref": "#/$defs/Location"}, + "preferences": {"anyOf": [{"type": "object", "properties": {"units": {"type": ["string", "null"]}}}, {"type": "null"}]}, + "attachments": {"type": "array", "items": {"anyOf": [{"type": "object"}, {"type": "null"}]}} + }, + "required": ["location"] + }` + parameters := mustGigaChatToolParameters(t, rawSchema) + before, err := schemas.MarshalSorted(parameters) + if err != nil { + t.Fatalf("failed to marshal original parameters: %v", err) + } + + chatRequest := testGigaChatChatToolRequest(t, "get_weather") + chatRequest.Params.Tools[0].Function.Parameters = parameters + gigaChatReq, err := ToGigaChatChatRequest(testBifrostContext(), chatRequest) + if err != nil { + t.Fatalf("ToGigaChatChatRequest returned error: %v", err) + } + after, err := schemas.MarshalSorted(parameters) + if err != nil { + t.Fatalf("failed to marshal original parameters after conversion: %v", err) + } + if string(before) != string(after) { + t.Fatalf("sanitizer mutated input schema:\nbefore=%s\nafter=%s", before, after) + } + + sanitized := mustGigaChatParametersMap(t, gigaChatReq.Functions[0].Parameters) + if _, exists := sanitized["$defs"]; exists { + t.Fatalf("sanitized schema still has $defs: %#v", sanitized) + } + properties := sanitized["properties"].(map[string]interface{}) + nickname := properties["nickname"].(map[string]interface{}) + if nickname["type"] != "string" { + t.Fatalf("nullable string was not sanitized: %#v", nickname) + } + if _, exists := nickname["nullable"]; exists { + t.Fatalf("nullable flag was not removed: %#v", nickname) + } + location := properties["location"].(map[string]interface{}) + locationProperties := location["properties"].(map[string]interface{}) + city := locationProperties["city"].(map[string]interface{}) + if city["type"] != "string" { + t.Fatalf("$ref optional string was not sanitized: %#v", city) + } + coords := locationProperties["coords"].(map[string]interface{}) + if coords["type"] != "object" || len(coords["properties"].(map[string]interface{})) != 0 { + t.Fatalf("nullable object without properties was not sanitized: %#v", coords) + } + preferences := properties["preferences"].(map[string]interface{}) + units := preferences["properties"].(map[string]interface{})["units"].(map[string]interface{}) + if preferences["type"] != "object" || units["type"] != "string" { + t.Fatalf("nested optional object was not sanitized: %#v", preferences) + } + attachments := properties["attachments"].(map[string]interface{}) + items := attachments["items"].(map[string]interface{}) + if items["type"] != "object" || len(items["properties"].(map[string]interface{})) != 0 { + t.Fatalf("array optional object item was not sanitized: %#v", items) + } + required := sanitized["required"].([]interface{}) + if len(required) != 1 || required[0] != "location" { + t.Fatalf("required fields changed unexpectedly: %#v", required) + } + + responsesRequest := testGigaChatResponsesToolRequest(t, "web_search") + responsesRequest.Params.Tools[0].ResponsesToolFunction.Parameters = parameters + gigaChatResponsesReq, err := ToGigaChatResponsesRequest(responsesRequest) + if err != nil { + t.Fatalf("ToGigaChatResponsesRequest returned error: %v", err) + } + specification := gigaChatResponsesReq.Tools[0].Functions.Specifications[0] + if specification.Name != "__bifrost_gigachat_user_web_search" { + t.Fatalf("reserved function name was not remapped: %q", specification.Name) + } + responsesSanitized := mustGigaChatParametersMap(t, specification.Parameters) + if _, exists := responsesSanitized["$defs"]; exists { + t.Fatalf("responses sanitized schema still has $defs: %#v", responsesSanitized) + } +} + +func TestGigaChatFunctionSchemaSanitizerRejectsAmbiguousUnions(t *testing.T) { + t.Parallel() + + parameters := mustGigaChatToolParameters(t, `{ + "type": "object", + "properties": { + "value": {"anyOf": [{"type": "string"}, {"type": "number"}, {"type": "null"}]} + } + }`) + _, err := sanitizeGigaChatFunctionSchema(parameters) + if err == nil || !strings.Contains(err.Error(), "multiple non-null branches") { + t.Fatalf("expected ambiguous union error, got %v", err) + } +} + +func TestGigaChatFunctionSchemaSanitizerRejectsNullOnlyAllOf(t *testing.T) { + t.Parallel() + + parameters := mustGigaChatToolParameters(t, `{ + "type": "object", + "properties": { + "value": {"allOf": [{"type": "null"}]} + } + }`) + _, err := sanitizeGigaChatFunctionSchema(parameters) + if err == nil || !strings.Contains(err.Error(), "$.properties.value") || !strings.Contains(err.Error(), "null-only schemas") { + t.Fatalf("expected null-only allOf error, got %v", err) + } +} + +func TestGigaChatFunctionSchemaSanitizerHandlesTopLevelNullableObject(t *testing.T) { + t.Parallel() + + sanitized, err := sanitizeGigaChatFunctionSchema(map[string]interface{}{ + "type": []interface{}{"object", "null"}, + }) + if err != nil { + t.Fatalf("sanitizeGigaChatFunctionSchema returned error: %v", err) + } + got := mustGigaChatParametersMap(t, sanitized) + if got["type"] != "object" { + t.Fatalf("top-level nullable object was not sanitized: %#v", got) + } + if len(got["properties"].(map[string]interface{})) != 0 { + t.Fatalf("top-level object properties mismatch: %#v", got["properties"]) + } +} + +func testGigaChatToolsResponsesRejectsUnsupportedPolicy(t *testing.T) { + t.Parallel() + + parallelToolCalls := true + strict := true + timeTool := schemas.ResponsesTool{ + Type: schemas.ResponsesToolTypeFunction, + Name: schemas.Ptr("get_time"), + Description: schemas.Ptr("Gets current time."), + ResponsesToolFunction: &schemas.ResponsesToolFunction{ + Parameters: mustGigaChatToolParameters(t, `{"type":"object","properties":{"timezone":{"type":"string"}},"required":["timezone"]}`), + }, + } + tests := []struct { + name string + mutate func(*schemas.BifrostResponsesRequest) + wantErr string + }{ + { + name: "InvalidJSONSchema", + mutate: func(request *schemas.BifrostResponsesRequest) { + request.Params.Tools[0].ResponsesToolFunction.Parameters = invalidGigaChatToolParameters() + }, + wantErr: "JSON schema is invalid", + }, + { + name: "GigaChatBuiltInName", + mutate: func(request *schemas.BifrostResponsesRequest) { + request.Params.Tools[0].Name = schemas.Ptr("text2image") + }, + wantErr: "built-in function", + }, + { + name: "OpenAIStrictMode", + mutate: func(request *schemas.BifrostResponsesRequest) { + request.Params.Tools[0].ResponsesToolFunction.Strict = &strict + }, + wantErr: "strict mode", + }, + { + name: "UnsupportedHostedTool", + mutate: func(request *schemas.BifrostResponsesRequest) { + request.Params.Tools = []schemas.ResponsesTool{{ + Type: schemas.ResponsesToolTypeFileSearch, + ResponsesToolFileSearch: &schemas.ResponsesToolFileSearch{ + VectorStoreIDs: []string{"vs_123"}, + }, + }} + }, + wantErr: "does not support tool type", + }, + { + name: "ParallelToolCalls", + mutate: func(request *schemas.BifrostResponsesRequest) { + request.Params.ParallelToolCalls = ¶llelToolCalls + }, + wantErr: "parallel_tool_calls", + }, + { + name: "RequiredToolChoiceWithMultipleTools", + mutate: func(request *schemas.BifrostResponsesRequest) { + request.Params.Tools = append(request.Params.Tools, timeTool) + request.Params.ToolChoice = &schemas.ResponsesToolChoice{ResponsesToolChoiceStr: schemas.Ptr("required")} + }, + wantErr: "cannot require an arbitrary", + }, + { + name: "AnyToolChoiceWithMultipleTools", + mutate: func(request *schemas.BifrostResponsesRequest) { + request.Params.Tools = append(request.Params.Tools, timeTool) + request.Params.ToolChoice = &schemas.ResponsesToolChoice{ResponsesToolChoiceStr: schemas.Ptr("any")} + }, + wantErr: "cannot require an arbitrary", + }, + { + name: "StructRequiredToolChoiceWithMultipleTools", + mutate: func(request *schemas.BifrostResponsesRequest) { + request.Params.Tools = append(request.Params.Tools, timeTool) + request.Params.ToolChoice = &schemas.ResponsesToolChoice{ResponsesToolChoiceStruct: &schemas.ResponsesToolChoiceStruct{ + Type: schemas.ResponsesToolChoiceTypeRequired, + }} + }, + wantErr: "cannot require an arbitrary", + }, + { + name: "AllowedToolsChoice", + mutate: func(request *schemas.BifrostResponsesRequest) { + request.Params.ToolChoice = &schemas.ResponsesToolChoice{ResponsesToolChoiceStruct: &schemas.ResponsesToolChoiceStruct{ + Type: schemas.ResponsesToolChoiceTypeAllowedTools, + Tools: []schemas.ResponsesToolChoiceAllowedToolDef{{ + Type: string(schemas.ResponsesToolTypeFunction), + Name: schemas.Ptr("get_weather"), + }}, + }} + }, + wantErr: "allowed tools set", + }, + { + name: "AutoToolChoiceWithoutFunctions", + mutate: func(request *schemas.BifrostResponsesRequest) { + request.Params.Tools = nil + request.Params.ToolChoice = &schemas.ResponsesToolChoice{ResponsesToolChoiceStr: schemas.Ptr("auto")} + }, + wantErr: "requires at least one", + }, + { + name: "UnknownForcedToolChoice", + mutate: func(request *schemas.BifrostResponsesRequest) { + request.Params.ToolChoice = &schemas.ResponsesToolChoice{ResponsesToolChoiceStruct: &schemas.ResponsesToolChoiceStruct{ + Type: schemas.ResponsesToolChoiceTypeFunction, + Name: schemas.Ptr("missing_tool"), + }} + }, + wantErr: "must match", + }, + { + name: "UnknownBuiltInToolChoice", + mutate: func(request *schemas.BifrostResponsesRequest) { + request.Params.ToolChoice = &schemas.ResponsesToolChoice{ResponsesToolChoiceStruct: &schemas.ResponsesToolChoiceStruct{ + Type: schemas.ResponsesToolChoiceTypeCodeInterpreter, + }} + }, + wantErr: "must match a declared", + }, + { + name: "ConflictingDuplicateFunction", + mutate: func(request *schemas.BifrostResponsesRequest) { + duplicate := request.Params.Tools[0] + duplicate.Description = schemas.Ptr("Different weather function.") + request.Params.Tools = append(request.Params.Tools, duplicate) + }, + wantErr: "different definition", + }, + { + name: "ExtraParamToolConfigBypass", + mutate: func(request *schemas.BifrostResponsesRequest) { + request.Params.ExtraParams = map[string]interface{}{"tool_config": map[string]interface{}{"mode": "auto"}} + }, + wantErr: "extra_params.tool_config", + }, + } + + for _, test := range tests { + test := test + t.Run(test.name, func(t *testing.T) { + t.Parallel() + + request := testGigaChatResponsesToolRequest(t, "get_weather") + test.mutate(request) + _, err := ToGigaChatResponsesRequest(request) + if err == nil || !strings.Contains(err.Error(), test.wantErr) { + t.Fatalf("expected error containing %q, got %v", test.wantErr, err) + } + }) + } +} + +func testGigaChatChatToolRequest(t *testing.T, toolName string) *schemas.BifrostChatRequest { + t.Helper() + + return &schemas.BifrostChatRequest{ + Model: "GigaChat", + Input: []schemas.ChatMessage{{ + Role: schemas.ChatMessageRoleUser, + Content: &schemas.ChatMessageContent{ContentStr: schemas.Ptr("Weather?")}, + }}, + Params: &schemas.ChatParameters{ + Tools: []schemas.ChatTool{testGigaChatChatFunctionTool(t, toolName)}, + }, + } +} + +func testGigaChatResponsesToolRequest(t *testing.T, toolName string) *schemas.BifrostResponsesRequest { + t.Helper() + + return &schemas.BifrostResponsesRequest{ + Model: "GigaChat-2", + Input: []schemas.ResponsesMessage{{ + Role: schemas.Ptr(schemas.ResponsesInputMessageRoleUser), + Content: &schemas.ResponsesMessageContent{ContentStr: schemas.Ptr("Weather?")}, + }}, + Params: &schemas.ResponsesParameters{ + Tools: []schemas.ResponsesTool{{ + Type: schemas.ResponsesToolTypeFunction, + Name: schemas.Ptr(toolName), + Description: schemas.Ptr("Gets current weather."), + ResponsesToolFunction: &schemas.ResponsesToolFunction{ + Parameters: mustGigaChatToolParameters(t, `{"type":"object","properties":{"city":{"type":"string"}},"required":["city"]}`), + }, + }}, + }, + } +} + +func testGigaChatChatFunctionTool(t *testing.T, toolName string) schemas.ChatTool { + t.Helper() + + return schemas.ChatTool{ + Type: schemas.ChatToolTypeFunction, + Function: &schemas.ChatToolFunction{ + Name: toolName, + Description: schemas.Ptr("Gets current weather."), + Parameters: mustGigaChatToolParameters(t, `{"type":"object","properties":{"city":{"type":"string"}},"required":["city"]}`), + }, + } +} + +func invalidGigaChatToolParameters() *schemas.ToolFunctionParameters { + return &schemas.ToolFunctionParameters{ + Type: "object", + AdditionalProperties: &schemas.AdditionalPropertiesStruct{}, + } +} + +func mustGigaChatParametersMap(t *testing.T, parameters *schemas.ToolFunctionParameters) map[string]interface{} { + t.Helper() + + raw, err := schemas.MarshalSorted(parameters) + if err != nil { + t.Fatalf("failed to marshal parameters: %v", err) + } + var out map[string]interface{} + if err := json.Unmarshal(raw, &out); err != nil { + t.Fatalf("failed to unmarshal parameters: %v", err) + } + return out +} diff --git a/core/providers/gigachat/types_test.go b/core/providers/gigachat/types_test.go new file mode 100644 index 00000000000..d7ae8ce2f27 --- /dev/null +++ b/core/providers/gigachat/types_test.go @@ -0,0 +1,94 @@ +package gigachat + +import ( + "encoding/json" + "testing" +) + +func TestGigaChatResponsesToolMarshalBuiltInTypes(t *testing.T) { + t.Parallel() + + stringPtr := func(value string) *string { + return &value + } + + tests := []struct { + name string + tool GigaChatResponsesTool + want string + }{ + { + name: "CodeInterpreter", + tool: GigaChatResponsesTool{ + CodeInterpreter: map[string]interface{}{"max_execution_time": 60}, + }, + want: `{"code_interpreter":{"max_execution_time":60}}`, + }, + { + name: "ImageGenerate", + tool: GigaChatResponsesTool{ + ImageGenerate: map[string]interface{}{"model": "Kandinsky", "width": 1024}, + }, + want: `{"image_generate":{"model":"Kandinsky","width":1024}}`, + }, + { + name: "ImageGenerateEmpty", + tool: GigaChatResponsesTool{ + ImageGenerate: map[string]interface{}{}, + }, + want: `{"image_generate":{}}`, + }, + { + name: "WebSearch", + tool: GigaChatResponsesTool{ + WebSearch: &GigaChatResponsesWebSearchTool{ + Type: stringPtr("search_plus"), + Indexes: []string{"public"}, + Flags: []string{"exact"}, + }, + }, + want: `{"web_search":{"type":"search_plus","indexes":["public"],"flags":["exact"]}}`, + }, + { + name: "URLContentExtraction", + tool: GigaChatResponsesTool{ + URLContentExtraction: map[string]interface{}{"mode": "summary"}, + }, + want: `{"url_content_extraction":{"mode":"summary"}}`, + }, + { + name: "Model3DGenerate", + tool: GigaChatResponsesTool{ + Model3DGenerate: map[string]interface{}{"format": "glb"}, + }, + want: `{"model_3d_generate":{"format":"glb"}}`, + }, + { + name: "Functions", + tool: GigaChatResponsesTool{ + Functions: &GigaChatResponsesFunctionsTool{ + Specifications: []GigaChatResponsesFunctionSpecification{{ + Name: "get_weather", + Parameters: mustGigaChatToolParameters(t, `{"type":"object","properties":{}}`), + }}, + }, + }, + want: `{"functions":{"specifications":[{"name":"get_weather","parameters":{"type":"object","properties":{}}}]}}`, + }, + } + + for _, test := range tests { + test := test + t.Run(test.name, func(t *testing.T) { + t.Parallel() + + raw, err := json.Marshal(test.tool) + if err != nil { + t.Fatalf("failed to marshal tool: %v", err) + } + if string(raw) != test.want { + t.Fatalf("tool JSON mismatch:\n got: %s\nwant: %s", raw, test.want) + } + }) + } +} diff --git a/core/providers/gigachat/utils_test.go b/core/providers/gigachat/utils_test.go new file mode 100644 index 00000000000..c4ae1a39222 --- /dev/null +++ b/core/providers/gigachat/utils_test.go @@ -0,0 +1,256 @@ +package gigachat + +import ( + "context" + "strings" + "testing" + + "github.com/maximhq/bifrost/core/schemas" +) + +func TestBuildGigaChatURL(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + baseURL string + apiVersion string + path string + want string + }{ + { + name: "default base appends v1 path", + apiVersion: gigaChatAPIVersionV1, + path: "/chat/completions", + want: "https://gigachat.devices.sberbank.ru/api/v1/chat/completions", + }, + { + name: "primary api root appends v1 path", + baseURL: "https://gigachat.devices.sberbank.ru/api/", + apiVersion: gigaChatAPIVersionV1, + path: "/chat/completions", + want: "https://gigachat.devices.sberbank.ru/api/v1/chat/completions", + }, + { + name: "primary v1 root keeps v1 path", + baseURL: "https://gigachat.devices.sberbank.ru/api/v1", + apiVersion: gigaChatAPIVersionV1, + path: "/chat/completions", + want: "https://gigachat.devices.sberbank.ru/api/v1/chat/completions", + }, + { + name: "primary v1 root switches to v2", + baseURL: "https://gigachat.devices.sberbank.ru/api/v1", + apiVersion: gigaChatAPIVersionV2, + path: "/chat/completions", + want: "https://gigachat.devices.sberbank.ru/api/v2/chat/completions", + }, + { + name: "primary v2 root switches to v1", + baseURL: "https://gigachat.devices.sberbank.ru/api/v2/", + apiVersion: gigaChatAPIVersionV1, + path: "/chat/completions", + want: "https://gigachat.devices.sberbank.ru/api/v1/chat/completions", + }, + { + name: "business api root appends v1 path", + baseURL: "https://api.giga.chat", + apiVersion: gigaChatAPIVersionV1, + path: "chat/completions", + want: "https://api.giga.chat/v1/chat/completions", + }, + { + name: "business v2 root keeps v2 path", + baseURL: "https://api.giga.chat/v2/", + apiVersion: gigaChatAPIVersionV2, + path: "/chat/completions", + want: "https://api.giga.chat/v2/chat/completions", + }, + { + name: "versioned path is not duplicated", + baseURL: "https://api.giga.chat", + apiVersion: gigaChatAPIVersionV1, + path: "/v1/chat/completions", + want: "https://api.giga.chat/v1/chat/completions", + }, + { + name: "path version is replaced by selected version", + baseURL: "https://api.giga.chat", + apiVersion: gigaChatAPIVersionV2, + path: "/v1/chat/completions", + want: "https://api.giga.chat/v2/chat/completions", + }, + { + name: "query string is preserved", + baseURL: "https://gigachat.devices.sberbank.ru/api", + apiVersion: gigaChatAPIVersionV1, + path: "/models?limit=100", + want: "https://gigachat.devices.sberbank.ru/api/v1/models?limit=100", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + got := buildGigaChatURL(tt.baseURL, tt.apiVersion, tt.path) + if got != tt.want { + t.Fatalf("buildGigaChatURL() = %q, want %q", got, tt.want) + } + }) + } +} + +func TestRedactGigaChatSensitiveText(t *testing.T) { + t.Parallel() + + input := `"access_token":"double-secret" "password": "spaced-secret" 'client_secret':'single-secret' credentials=plain-secret authorization=assignment-secret user=auth-secret-user Bearer header-secret -----BEGIN PRIVATE KEY----- +private-secret +-----END PRIVATE KEY-----` + + got := redactGigaChatSensitiveText(input) + for _, secret := range []string{"double-secret", "spaced-secret", "single-secret", "plain-secret", "assignment-secret", "auth-secret-user", "header-secret", "private-secret"} { + if strings.Contains(got, secret) { + t.Fatalf("redacted text leaked %q in %s", secret, got) + } + } + for _, want := range []string{ + `"access_token":""`, + `"password": ""`, + `'client_secret':''`, + `credentials=`, + `authorization=`, + `user=`, + `Bearer `, + ``, + } { + if !strings.Contains(got, want) { + t.Fatalf("redacted text missing %q in %s", want, got) + } + } + + benign := redactGigaChatSensitiveText("profile user=visible-user") + if !strings.Contains(benign, "user=visible-user") { + t.Fatalf("benign user assignment should not be redacted: %s", benign) + } +} + +func TestRedactGigaChatRawPayloadNarrowsUserField(t *testing.T) { + t.Parallel() + + payload := []byte(`{ + "user": "ordinary-user", + "message": "keep this", + "gigachat_key_config": { + "user": {"value": "secret-user"}, + "password": {"value": "secret-password"}, + "credentials": {"value": "secret-credentials"}, + "access_token": {"value": "secret-token"} + } + }`) + + redacted := string(redactGigaChatRawPayload(payload)) + for _, secret := range []string{"secret-user", "secret-password", "secret-credentials", "secret-token"} { + if strings.Contains(redacted, secret) { + t.Fatalf("redacted payload leaked %q in %s", secret, redacted) + } + } + if !strings.Contains(redacted, `"user":"ordinary-user"`) { + t.Fatalf("benign top-level user field should be preserved: %s", redacted) + } + for _, want := range []string{`"user":""`, `"password":""`, `"credentials":""`, `"access_token":""`} { + if !strings.Contains(redacted, want) { + t.Fatalf("redacted payload missing %q in %s", want, redacted) + } + } +} + +func TestResolveGigaChatURLs(t *testing.T) { + t.Parallel() + + t.Run("auth URL defaults", func(t *testing.T) { + t.Parallel() + + if got := resolveAuthURL(schemas.Key{}); got != gigaChatDefaultAuthURL { + t.Fatalf("resolveAuthURL() = %q, want %q", got, gigaChatDefaultAuthURL) + } + }) + + t.Run("auth URL uses key override", func(t *testing.T) { + t.Parallel() + + key := schemas.Key{ + GigaChatKeyConfig: &schemas.GigaChatKeyConfig{ + AuthURL: "https://auth.example.com/oauth/", + }, + } + if got := resolveAuthURL(key); got != "https://auth.example.com/oauth" { + t.Fatalf("resolveAuthURL() = %q", got) + } + }) + + t.Run("base URL defaults", func(t *testing.T) { + t.Parallel() + + if got := resolveBaseURL(schemas.Key{}, schemas.NetworkConfig{}); got != gigaChatDefaultBaseURL { + t.Fatalf("resolveBaseURL() = %q, want %q", got, gigaChatDefaultBaseURL) + } + }) + + t.Run("base URL uses provider network config", func(t *testing.T) { + t.Parallel() + + networkConfig := schemas.NetworkConfig{ + BaseURL: "https://api.giga.chat/v1/", + } + if got := resolveBaseURL(schemas.Key{}, networkConfig); got != "https://api.giga.chat/v1" { + t.Fatalf("resolveBaseURL() = %q", got) + } + }) + + t.Run("base URL uses key override before provider network config", func(t *testing.T) { + t.Parallel() + + key := schemas.Key{ + GigaChatKeyConfig: &schemas.GigaChatKeyConfig{ + BaseURL: "https://api.giga.chat/v2/", + }, + } + networkConfig := schemas.NetworkConfig{ + BaseURL: "https://gigachat.devices.sberbank.ru/api/v1", + } + if got := resolveBaseURL(key, networkConfig); got != "https://api.giga.chat/v2" { + t.Fatalf("resolveBaseURL() = %q", got) + } + }) +} + +func TestBuildGigaChatRequestURL(t *testing.T) { + t.Parallel() + + t.Run("context path override is version-normalized", func(t *testing.T) { + t.Parallel() + + ctx := schemas.NewBifrostContext(context.Background(), schemas.NoDeadline) + ctx.SetValue(schemas.BifrostContextKeyURLPath, "/v1/custom/chat") + + got := buildGigaChatRequestURL(ctx, "https://api.giga.chat/v1", gigaChatAPIVersionV2, "/chat/completions", nil, schemas.ChatCompletionRequest) + want := "https://api.giga.chat/v2/custom/chat" + if got != want { + t.Fatalf("buildGigaChatRequestURL() = %q, want %q", got, want) + } + }) + + t.Run("absolute context URL override is preserved", func(t *testing.T) { + t.Parallel() + + ctx := schemas.NewBifrostContext(context.Background(), schemas.NoDeadline) + ctx.SetValue(schemas.BifrostContextKeyURLPath, "https://proxy.example.com/gigachat") + + got := buildGigaChatRequestURL(ctx, "https://api.giga.chat", gigaChatAPIVersionV1, "/chat/completions", nil, schemas.ChatCompletionRequest) + want := "https://proxy.example.com/gigachat" + if got != want { + t.Fatalf("buildGigaChatRequestURL() = %q, want %q", got, want) + } + }) +} diff --git a/framework/configstore/clientconfig_redaction_test.go b/framework/configstore/clientconfig_redaction_test.go index 204bac20748..8328c1c7f49 100644 --- a/framework/configstore/clientconfig_redaction_test.go +++ b/framework/configstore/clientconfig_redaction_test.go @@ -163,6 +163,7 @@ func TestProviderConfig_Redacted_FullJSONHasNoLeakedEnvSecrets(t *testing.T) { t.Setenv("LEAK_TEST_VERTEX_PROJECT", "leaked-vertex-project-id") t.Setenv("LEAK_TEST_BEDROCK_ACCESS", "AKIAIOSFODNN7LEAKED1") t.Setenv("LEAK_TEST_OPENAI_KEY", "sk-leaked-openai-key-1234567890") + t.Setenv("LEAK_TEST_GIGACHAT_CREDENTIALS", "leaked-gigachat-credentials") config := ProviderConfig{ Keys: []schemas.Key{ @@ -197,6 +198,17 @@ func TestProviderConfig_Redacted_FullJSONHasNoLeakedEnvSecrets(t *testing.T) { SecretKey: schemas.SecretVar{Val: ""}, }, }, + { + ID: "gigachat-k", + Name: "gigachat", + Value: schemas.SecretVar{Val: ""}, + GigaChatKeyConfig: &schemas.GigaChatKeyConfig{ + Credentials: schemas.NewSecretVar("env.LEAK_TEST_GIGACHAT_CREDENTIALS"), + CertFile: "/secure/client.pem", + KeyFile: "/secure/client.key", + CABundleFile: "/secure/ca.pem", + }, + }, }, } @@ -210,6 +222,10 @@ func TestProviderConfig_Redacted_FullJSONHasNoLeakedEnvSecrets(t *testing.T) { "leaked-vertex-project-id", "AKIAIOSFODNN7LEAKED1", "sk-leaked-openai-key-1234567890", + "leaked-gigachat-credentials", + "/secure/client.pem", + "/secure/client.key", + "/secure/ca.pem", } for _, secret := range leakedSecrets { assert.False(t, strings.Contains(jsonStr, secret), @@ -222,6 +238,7 @@ func TestProviderConfig_Redacted_FullJSONHasNoLeakedEnvSecrets(t *testing.T) { "env.LEAK_TEST_AZURE_ENDPOINT", "env.LEAK_TEST_VERTEX_PROJECT", "env.LEAK_TEST_BEDROCK_ACCESS", + "env.LEAK_TEST_GIGACHAT_CREDENTIALS", } for _, ref := range expectedRefs { assert.True(t, strings.Contains(jsonStr, ref), diff --git a/framework/configstore/migrations_test.go b/framework/configstore/migrations_test.go index 8398626e1ce..8bb9bdc4589 100644 --- a/framework/configstore/migrations_test.go +++ b/framework/configstore/migrations_test.go @@ -1176,6 +1176,7 @@ func TestTriggerMigrations_FreshDB(t *testing.T) { for _, table := range criticalTables { assert.True(t, migrator.HasTable(table), "table should exist: %T", table) } + assert.True(t, migrator.HasColumn(&tables.TableKey{}, "gigachat_key_config_json"), "GigaChat key config column should exist") assert.True(t, migrator.HasColumn(&tables.TableModelPricing{}, "is_deprecated"), "model pricing is_deprecated column should exist") } @@ -2813,6 +2814,125 @@ func TestMigrationAddDualCredentialConflictBehaviorColumn(t *testing.T) { "re-running the migration should be idempotent") } +func TestGenerateKeyHashGigaChatKeyConfigFields(t *testing.T) { + newKey := func() schemas.Key { + return schemas.Key{ + Name: "gigachat-key", + Value: *schemas.NewSecretVar("unused"), + Weight: 1, + GigaChatKeyConfig: &schemas.GigaChatKeyConfig{ + Credentials: schemas.NewSecretVar("credentials"), + Scope: schemas.DefaultGigaChatScope, + User: schemas.NewSecretVar("user"), + Password: schemas.NewSecretVar("password"), + AccessToken: schemas.NewSecretVar("access-token"), + AuthURL: "https://auth.example.com", + BaseURL: "https://api.example.com", + CertFile: "/tls/client.pem", + KeyFile: "/tls/client.key", + CABundleFile: "/tls/ca.pem", + }, + } + } + + baseHash, err := GenerateKeyHash(newKey()) + require.NoError(t, err) + + tests := []struct { + name string + mutate func(*schemas.GigaChatKeyConfig) + }{ + {"credentials", func(config *schemas.GigaChatKeyConfig) { config.Credentials = schemas.NewSecretVar("new-credentials") }}, + {"scope", func(config *schemas.GigaChatKeyConfig) { config.Scope = "GIGACHAT_API_CORP" }}, + {"user", func(config *schemas.GigaChatKeyConfig) { config.User = schemas.NewSecretVar("new-user") }}, + {"password", func(config *schemas.GigaChatKeyConfig) { config.Password = schemas.NewSecretVar("new-password") }}, + {"access token", func(config *schemas.GigaChatKeyConfig) { config.AccessToken = schemas.NewSecretVar("new-access-token") }}, + {"auth URL", func(config *schemas.GigaChatKeyConfig) { config.AuthURL = "https://new-auth.example.com" }}, + {"base URL", func(config *schemas.GigaChatKeyConfig) { config.BaseURL = "https://new-api.example.com" }}, + {"certificate file", func(config *schemas.GigaChatKeyConfig) { config.CertFile = "/tls/new-client.pem" }}, + {"key file", func(config *schemas.GigaChatKeyConfig) { config.KeyFile = "/tls/new-client.key" }}, + {"CA bundle file", func(config *schemas.GigaChatKeyConfig) { config.CABundleFile = "/tls/new-ca.pem" }}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + key := newKey() + tt.mutate(key.GigaChatKeyConfig) + changedHash, err := GenerateKeyHash(key) + require.NoError(t, err) + require.NotEqual(t, baseHash, changedHash) + }) + } +} + +func TestMigrationAddGigaChatKeyConfigColumnBackfillsHashes(t *testing.T) { + db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{ + Logger: logger.Default.LogMode(logger.Silent), + }) + require.NoError(t, err) + + require.NoError(t, db.Exec(`CREATE TABLE migrations (id VARCHAR(255) PRIMARY KEY)`).Error) + require.NoError(t, db.Exec(` + CREATE TABLE config_keys ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + name VARCHAR(255) NOT NULL UNIQUE, + provider_id INTEGER NOT NULL, + provider VARCHAR(50), + key_id VARCHAR(255) NOT NULL UNIQUE, + value TEXT NOT NULL, + models_json TEXT, + blacklisted_models_json TEXT, + weight REAL, + enabled BOOLEAN DEFAULT true, + created_at DATETIME NOT NULL, + updated_at DATETIME NOT NULL, + config_hash VARCHAR(255), + use_for_batch_api BOOLEAN DEFAULT false, + use_anthropic_endpoints BOOLEAN DEFAULT false, + status VARCHAR(50) DEFAULT 'unknown', + description TEXT, + encryption_status VARCHAR(20) DEFAULT 'plain_text' + ) + `).Error) + require.False(t, db.Migrator().HasColumn(&tables.TableKey{}, "gigachat_key_config_json"), + "precondition: legacy config_keys schema must not contain gigachat_key_config_json") + + now := time.Now() + require.NoError(t, db.Exec(` + INSERT INTO config_keys ( + name, provider_id, provider, key_id, value, models_json, weight, enabled, + created_at, updated_at, config_hash, encryption_status + ) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + `, + "gigachat-key", 1, string(schemas.GigaChat), "gigachat-key-id", "unused", `["*"]`, 1.0, true, + now, now, "stale-gigachat-hash", tables.EncryptionStatusPlainText, + ).Error) + require.NoError(t, db.Exec(` + INSERT INTO config_keys ( + name, provider_id, provider, key_id, value, models_json, weight, enabled, + created_at, updated_at, config_hash, encryption_status + ) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + `, + "openai-key", 2, string(schemas.OpenAI), "openai-key-id", "sk-test", `["*"]`, 1.0, true, + now, now, "untouched-openai-hash", tables.EncryptionStatusPlainText, + ).Error) + + require.NoError(t, migrationAddGigaChatKeyConfigColumn(context.Background(), db, testMigrationLogger)) + require.True(t, db.Migrator().HasColumn(&tables.TableKey{}, "gigachat_key_config_json"), + "migration should add gigachat_key_config_json to the legacy schema") + + var migrated tables.TableKey + require.NoError(t, db.Where("key_id = ?", "gigachat-key-id").First(&migrated).Error) + expectedHash, err := GenerateKeyHash(schemaKeyFromTableKey(migrated)) + require.NoError(t, err) + require.Equal(t, expectedHash, migrated.ConfigHash) + require.NotEqual(t, "stale-gigachat-hash", migrated.ConfigHash) + + var untouched tables.TableKey + require.NoError(t, db.Where("key_id = ?", "openai-key-id").First(&untouched).Error) + require.Equal(t, "untouched-openai-hash", untouched.ConfigHash) +} + // TestMigrationAddBudgetOverrideColumns verifies legacy budgets receive inactive override defaults. func TestMigrationAddBudgetOverrideColumns(t *testing.T) { db := setupTestDB(t) diff --git a/framework/configstore/tables/encryption_test.go b/framework/configstore/tables/encryption_test.go index 50172df2c66..8903fac41ba 100644 --- a/framework/configstore/tables/encryption_test.go +++ b/framework/configstore/tables/encryption_test.go @@ -1343,6 +1343,15 @@ func TestTableKey_AllProviderConfigs_EncryptDecrypt(t *testing.T) { Region: schemas.NewSecretVar("eu-west-1"), ARN: schemas.NewSecretVar("arn:aws:bedrock:eu-west-1:123:role"), }, + GigaChatKeyConfig: &schemas.GigaChatKeyConfig{ + Credentials: schemas.NewSecretVar("gigachat-credentials"), + Scope: schemas.DefaultGigaChatScope, + BaseURL: "https://api.giga.chat", + AuthURL: "https://ngw.devices.sberbank.ru:9443/api/v2/oauth", + CertFile: "/secure/client.pem", + KeyFile: "/secure/client.key", + CABundleFile: "/secure/ca.pem", + }, BedrockMantleKeyConfig: &schemas.BedrockMantleKeyConfig{ AccessKey: *schemas.NewSecretVar("AKIA-MANTLE"), SecretKey: *schemas.NewSecretVar("wJalr-MANTLE"), @@ -1363,6 +1372,18 @@ func TestTableKey_AllProviderConfigs_EncryptDecrypt(t *testing.T) { assert.NotEqual(t, "us-central1", raw["vertex_region"]) assert.NotEqual(t, "eu-west-1", raw["bedrock_region"]) assert.NotEqual(t, "arn:aws:bedrock:eu-west-1:123:role", raw["bedrock_arn"]) + rawGigaChatVal := raw["gigachat_key_config_json"] + require.NotNil(t, rawGigaChatVal, "gigachat_key_config_json should be present in raw row") + var rawGigaChatStr string + switch v := rawGigaChatVal.(type) { + case string: + rawGigaChatStr = v + case []byte: + rawGigaChatStr = string(v) + } + require.NotEmpty(t, rawGigaChatStr, "gigachat_key_config_json should not be empty") + assert.NotContains(t, rawGigaChatStr, "gigachat-credentials") + assert.NotContains(t, rawGigaChatStr, "/secure/client.pem") assert.NotEqual(t, "proj_mantle456", raw["bedrock_mantle_project_id"]) rawAliasesVal2 := raw["aliases_json"] require.NotNil(t, rawAliasesVal2, "aliases_json should be present in raw row") @@ -1405,6 +1426,16 @@ func TestTableKey_AllProviderConfigs_EncryptDecrypt(t *testing.T) { require.NotNil(t, found.BedrockKeyConfig.ARN) assert.Equal(t, "arn:aws:bedrock:eu-west-1:123:role", found.BedrockKeyConfig.ARN.GetValue()) + require.NotNil(t, found.GigaChatKeyConfig) + require.NotNil(t, found.GigaChatKeyConfig.Credentials) + assert.Equal(t, "gigachat-credentials", found.GigaChatKeyConfig.Credentials.GetValue()) + assert.Equal(t, schemas.DefaultGigaChatScope, found.GigaChatKeyConfig.Scope) + assert.Equal(t, "https://api.giga.chat", found.GigaChatKeyConfig.BaseURL) + assert.Equal(t, "https://ngw.devices.sberbank.ru:9443/api/v2/oauth", found.GigaChatKeyConfig.AuthURL) + assert.Equal(t, "/secure/client.pem", found.GigaChatKeyConfig.CertFile) + assert.Equal(t, "/secure/client.key", found.GigaChatKeyConfig.KeyFile) + assert.Equal(t, "/secure/ca.pem", found.GigaChatKeyConfig.CABundleFile) + require.NotNil(t, found.BedrockMantleKeyConfig) assert.Equal(t, "AKIA-MANTLE", found.BedrockMantleKeyConfig.AccessKey.GetValue()) assert.Equal(t, "wJalr-MANTLE", found.BedrockMantleKeyConfig.SecretKey.GetValue()) @@ -2243,4 +2274,4 @@ func TestTableKey_AliasesJSON_LegacyInputRoundTrip(t *testing.T) { assert.Equal(t, "gpt-4o-deployment", got.ModelID) assert.Nil(t, got.ModelName) assert.Nil(t, got.ModelFamily) -} \ No newline at end of file +} diff --git a/transports/bifrost-http/handlers/provider_keys_test.go b/transports/bifrost-http/handlers/provider_keys_test.go index c0bfdc22d7f..e6f2fc528b4 100644 --- a/transports/bifrost-http/handlers/provider_keys_test.go +++ b/transports/bifrost-http/handlers/provider_keys_test.go @@ -149,6 +149,13 @@ func TestMergeUpdatedKey_ProviderConfigMaskedPreviews(t *testing.T) { VLLMKeyConfig: &schemas.VLLMKeyConfig{URL: secret("https://current.vllm.example.com")}, OllamaKeyConfig: &schemas.OllamaKeyConfig{URL: secret("https://current.ollama.example.com")}, SGLKeyConfig: &schemas.SGLKeyConfig{URL: secret("https://current.sgl.example.com")}, + GigaChatKeyConfig: &schemas.GigaChatKeyConfig{ + Credentials: secretPtr("gigachat-credentials-current"), + AccessToken: secretPtr("gigachat-access-token-current"), + CertFile: "/secure/current-client.pem", + KeyFile: "/secure/current-client.key", + CABundleFile: "/secure/current-ca.pem", + }, } update := schemas.Key{ AzureKeyConfig: &schemas.AzureKeyConfig{ @@ -168,6 +175,13 @@ func TestMergeUpdatedKey_ProviderConfigMaskedPreviews(t *testing.T) { VLLMKeyConfig: &schemas.VLLMKeyConfig{URL: staleMask("vllm", "0007")}, OllamaKeyConfig: &schemas.OllamaKeyConfig{URL: staleMask("olla", "0008")}, SGLKeyConfig: &schemas.SGLKeyConfig{URL: staleMask("sgla", "0009")}, + GigaChatKeyConfig: &schemas.GigaChatKeyConfig{ + Credentials: secretPtr(""), + AccessToken: secretPtr(staleMaskValue("giga", "0010")), + CertFile: "", + KeyFile: "", + CABundleFile: "", + }, } merged := merge(oldRaw, update) @@ -185,6 +199,11 @@ func TestMergeUpdatedKey_ProviderConfigMaskedPreviews(t *testing.T) { {"vllm url", merged.VLLMKeyConfig.URL.GetValue(), oldRaw.VLLMKeyConfig.URL.GetValue()}, {"ollama url", merged.OllamaKeyConfig.URL.GetValue(), oldRaw.OllamaKeyConfig.URL.GetValue()}, {"sgl url", merged.SGLKeyConfig.URL.GetValue(), oldRaw.SGLKeyConfig.URL.GetValue()}, + {"gigachat credentials", merged.GigaChatKeyConfig.Credentials.GetValue(), oldRaw.GigaChatKeyConfig.Credentials.GetValue()}, + {"gigachat access token", merged.GigaChatKeyConfig.AccessToken.GetValue(), oldRaw.GigaChatKeyConfig.AccessToken.GetValue()}, + {"gigachat cert file", merged.GigaChatKeyConfig.CertFile, oldRaw.GigaChatKeyConfig.CertFile}, + {"gigachat key file", merged.GigaChatKeyConfig.KeyFile, oldRaw.GigaChatKeyConfig.KeyFile}, + {"gigachat CA bundle file", merged.GigaChatKeyConfig.CABundleFile, oldRaw.GigaChatKeyConfig.CABundleFile}, } for _, check := range checks { if check.got != check.want { @@ -197,6 +216,16 @@ func TestMergeUpdatedKey_ProviderConfigMaskedPreviews(t *testing.T) { if !merged.VLLMKeyConfig.URL.IsFromEnv() || merged.VLLMKeyConfig.URL.GetRawRef() != "env.NEW_VLLM_URL" { t.Fatalf("expected nested env ref applied, got ref=%q", merged.VLLMKeyConfig.URL.GetRawRef()) } + + update.GigaChatKeyConfig.Credentials = secretPtr("env.NEW_GIGACHAT_CREDENTIALS") + update.GigaChatKeyConfig.CertFile = "/secure/new-client.pem" + merged = merge(oldRaw, update) + if !merged.GigaChatKeyConfig.Credentials.IsFromEnv() || merged.GigaChatKeyConfig.Credentials.GetRawRef() != "env.NEW_GIGACHAT_CREDENTIALS" { + t.Fatalf("expected GigaChat env ref applied, got ref=%q", merged.GigaChatKeyConfig.Credentials.GetRawRef()) + } + if merged.GigaChatKeyConfig.CertFile != "/secure/new-client.pem" { + t.Fatalf("expected new GigaChat cert file applied, got %q", merged.GigaChatKeyConfig.CertFile) + } } func TestMergeUpdatedKey_RejectsMaskWithoutStoredCounterpart(t *testing.T) { @@ -233,6 +262,20 @@ func TestMergeUpdatedKey_RejectsMaskWithoutStoredCounterpart(t *testing.T) { update: schemas.Key{VLLMKeyConfig: &schemas.VLLMKeyConfig{URL: mask}}, wantErr: "vllm_key_config.url", }, + { + name: "missing GigaChat config section", + oldRaw: schemas.Key{}, + update: schemas.Key{GigaChatKeyConfig: &schemas.GigaChatKeyConfig{Credentials: schemas.NewSecretVar("")}}, + wantErr: "gigachat_key_config.credentials", + }, + { + name: "missing GigaChat TLS file", + oldRaw: schemas.Key{GigaChatKeyConfig: &schemas.GigaChatKeyConfig{ + CertFile: "/secure/current-client.pem", + }}, + update: schemas.Key{GigaChatKeyConfig: &schemas.GigaChatKeyConfig{KeyFile: ""}}, + wantErr: "gigachat_key_config.key_file", + }, } for _, tt := range tests { diff --git a/transports/bifrost-http/handlers/providers_test.go b/transports/bifrost-http/handlers/providers_test.go index b21e5b1c14e..2f138a03786 100644 --- a/transports/bifrost-http/handlers/providers_test.go +++ b/transports/bifrost-http/handlers/providers_test.go @@ -1664,3 +1664,73 @@ func TestListModels_KeyBlacklistIsCaseInsensitive(t *testing.T) { } } } + +func TestValidateProviderKeyURL_GigaChat(t *testing.T) { + tests := []struct { + name string + key schemas.Key + wantErr bool + }{ + { + name: "credentials config without key value is valid", + key: schemas.Key{ + GigaChatKeyConfig: &schemas.GigaChatKeyConfig{ + Credentials: schemas.NewSecretVar("env.GIGACHAT_CREDENTIALS"), + }, + }, + }, + { + name: "plain key value is valid", + key: schemas.Key{ + Value: *schemas.NewSecretVar("legacy-api-key"), + }, + }, + { + name: "client certificate config without key value is valid", + key: schemas.Key{ + GigaChatKeyConfig: &schemas.GigaChatKeyConfig{ + CertFile: "/secure/client.pem", + KeyFile: "/secure/client.key", + }, + }, + }, + { + name: "missing auth material is invalid", + key: schemas.Key{}, + wantErr: true, + }, + { + name: "partial user password config is invalid", + key: schemas.Key{ + GigaChatKeyConfig: &schemas.GigaChatKeyConfig{ + User: schemas.NewSecretVar("env.GIGACHAT_USER"), + }, + }, + wantErr: true, + }, + { + name: "partial certificate config is invalid", + key: schemas.Key{ + GigaChatKeyConfig: &schemas.GigaChatKeyConfig{ + CertFile: "/tmp/client.crt", + }, + }, + wantErr: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + err := validateProviderKeyURL(schemas.GigaChat, tt.key) + if tt.wantErr && err == nil { + t.Fatalf("expected error") + } + if !tt.wantErr && err != nil { + t.Fatalf("unexpected error: %v", err) + } + if tt.key.GigaChatKeyConfig != nil && tt.key.GigaChatKeyConfig.Scope == "" { + t.Fatalf("expected default GigaChat scope to be applied") + } + }) + } +} diff --git a/transports/bifrost-http/lib/validator_test.go b/transports/bifrost-http/lib/validator_test.go index a3966a43315..d131a247b56 100644 --- a/transports/bifrost-http/lib/validator_test.go +++ b/transports/bifrost-http/lib/validator_test.go @@ -1,6 +1,7 @@ package lib import ( + "fmt" "net/http" "net/http/httptest" "os" @@ -1584,6 +1585,226 @@ func TestValidateConfigSchema_BedrockKeyConfig_MissingRegion(t *testing.T) { } } +// ============================================================================= +// GigaChat Key Config Required Field Pair Tests +// Note: GigaChat provider uses a special key schema that extends base_key +// ============================================================================= + +func TestValidateConfigSchema_GigaChatKeyConfig_ValidCredentials(t *testing.T) { + validConfig := `{ + "providers": { + "gigachat": { + "keys": [ + { + "name": "gigachat-key", + "weight": 1.0, + "models": ["*"], + "gigachat_key_config": { + "credentials": "env.GIGACHAT_CREDENTIALS", + "scope": "GIGACHAT_API_PERS", + "base_url": "https://api.giga.chat" + } + } + ] + } + } + }` + + err := ValidateConfigSchema([]byte(validConfig), loadLocalSchema(t)) + if err != nil { + t.Errorf("expected valid GigaChat key config to pass validation, got: %v", err) + } +} + +func TestValidateConfigSchema_GigaChatKeyConfig_ValidAccessToken(t *testing.T) { + validConfig := `{ + "providers": { + "gigachat": { + "keys": [ + { + "name": "gigachat-key", + "weight": 1.0, + "models": ["*"], + "gigachat_key_config": { + "access_token": "env.GIGACHAT_ACCESS_TOKEN" + } + } + ] + } + } + }` + + err := ValidateConfigSchema([]byte(validConfig), loadLocalSchema(t)) + if err != nil { + t.Errorf("expected valid GigaChat access token config to pass validation, got: %v", err) + } +} + +func TestValidateConfigSchema_GigaChatKeyConfig_ValidUserPassword(t *testing.T) { + validConfig := `{ + "providers": { + "gigachat": { + "keys": [ + { + "name": "gigachat-key", + "weight": 1.0, + "models": ["*"], + "gigachat_key_config": { + "user": "env.GIGACHAT_USER", + "password": "env.GIGACHAT_PASSWORD" + } + } + ] + } + } + }` + + err := ValidateConfigSchema([]byte(validConfig), loadLocalSchema(t)) + if err != nil { + t.Errorf("expected valid GigaChat user/password config to pass validation, got: %v", err) + } +} + +func TestValidateConfigSchema_GigaChatKeyConfig_TLSOnlyValid(t *testing.T) { + validConfig := `{ + "providers": { + "gigachat": { + "keys": [ + { + "name": "gigachat-key", + "weight": 1.0, + "models": ["*"], + "gigachat_key_config": { + "cert_file": "/secure/client.pem", + "key_file": "/secure/client.key" + } + } + ] + } + } + }` + + err := ValidateConfigSchema([]byte(validConfig), loadLocalSchema(t)) + if err != nil { + t.Errorf("expected TLS-only GigaChat key config to pass validation, got: %v", err) + } +} + +func TestValidateConfigSchema_GigaChatKeyConfig_TLSWithCredentialsValid(t *testing.T) { + validConfig := `{ + "providers": { + "gigachat": { + "keys": [ + { + "name": "gigachat-key", + "weight": 1.0, + "models": ["*"], + "gigachat_key_config": { + "credentials": "env.GIGACHAT_CREDENTIALS", + "cert_file": "/secure/client.pem", + "key_file": "/secure/client.key", + "ca_bundle_file": "/secure/ca.pem" + } + } + ] + } + } + }` + + err := ValidateConfigSchema([]byte(validConfig), loadLocalSchema(t)) + if err != nil { + t.Errorf("expected GigaChat credentials plus TLS config to pass validation, got: %v", err) + } +} + +func TestValidateConfigSchema_GigaChatKeyConfig_MissingPassword(t *testing.T) { + invalidConfig := `{ + "providers": { + "gigachat": { + "keys": [ + { + "name": "gigachat-key", + "weight": 1.0, + "gigachat_key_config": { + "user": "env.GIGACHAT_USER" + } + } + ] + } + } + }` + + err := ValidateConfigSchema([]byte(invalidConfig), loadLocalSchema(t)) + if err == nil { + t.Error("expected config missing 'password' in GigaChat key config to fail validation") + } +} + +func TestValidateConfigSchema_GigaChatKeyConfig_MissingKeyFile(t *testing.T) { + invalidConfig := `{ + "providers": { + "gigachat": { + "keys": [ + { + "name": "gigachat-key", + "weight": 1.0, + "gigachat_key_config": { + "cert_file": "/secure/client.pem" + } + } + ] + } + } + }` + + err := ValidateConfigSchema([]byte(invalidConfig), loadLocalSchema(t)) + if err == nil { + t.Error("expected config missing 'key_file' in GigaChat key config to fail validation") + } +} + +func TestValidateConfigSchema_GigaChatKeyConfig_RejectsBlankAuthMaterial(t *testing.T) { + tests := []struct { + name string + keyJSON string + }{ + {name: "empty value", keyJSON: `"value": ""`}, + {name: "whitespace value", keyJSON: `"value": " "`}, + {name: "empty credentials", keyJSON: `"gigachat_key_config": {"credentials": ""}`}, + {name: "whitespace credentials", keyJSON: `"gigachat_key_config": {"credentials": " "}`}, + {name: "empty access token", keyJSON: `"gigachat_key_config": {"access_token": ""}`}, + {name: "whitespace access token", keyJSON: `"gigachat_key_config": {"access_token": " "}`}, + {name: "empty user", keyJSON: `"gigachat_key_config": {"user": "", "password": "secret"}`}, + {name: "whitespace user", keyJSON: `"gigachat_key_config": {"user": " ", "password": "secret"}`}, + {name: "empty password", keyJSON: `"gigachat_key_config": {"user": "user", "password": ""}`}, + {name: "whitespace password", keyJSON: `"gigachat_key_config": {"user": "user", "password": " "}`}, + {name: "empty cert file", keyJSON: `"gigachat_key_config": {"cert_file": "", "key_file": "/secure/client.key"}`}, + {name: "whitespace cert file", keyJSON: `"gigachat_key_config": {"cert_file": " ", "key_file": "/secure/client.key"}`}, + {name: "empty key file", keyJSON: `"gigachat_key_config": {"cert_file": "/secure/client.pem", "key_file": ""}`}, + {name: "whitespace key file", keyJSON: `"gigachat_key_config": {"cert_file": "/secure/client.pem", "key_file": " "}`}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + invalidConfig := fmt.Sprintf(`{ + "providers": { + "gigachat": { + "keys": [{ + "name": "gigachat-key", + "weight": 1.0, + %s + }] + } + } + }`, tt.keyJSON) + + if err := ValidateConfigSchema([]byte(invalidConfig), loadLocalSchema(t)); err == nil { + t.Fatalf("expected blank GigaChat auth material to fail validation: %s", invalidConfig) + } + }) + } +} + // ============================================================================= // Guardrails Config Tests // Note: Guardrails is an enterprise feature. The guardrails_config schema From 0f34d852c5223853b5a5f348b130e92e9138b060 Mon Sep 17 00:00:00 2001 From: krakenalt Date: Thu, 20 Aug 2026 13:49:26 +0300 Subject: [PATCH 4/6] [feat]: add GigaChat provider UI Expose GigaChat credentials, endpoints, provider metadata, and icons in the dashboard. --- .../fragments/apiKeysFormFragment.tsx | 287 +++++++++++++++++- .../views/modelProviderKeysTableView.tsx | 17 +- .../providers/views/providerKeyForm.tsx | 31 +- ui/lib/constants/config.ts | 2 + ui/lib/constants/icons.tsx | 14 + ui/lib/constants/logs.ts | 5 +- ui/lib/types/config.ts | 27 ++ ui/lib/types/schemas.ts | 66 ++++ ui/public/images/gigachat.svg | 30 ++ 9 files changed, 460 insertions(+), 19 deletions(-) create mode 100644 ui/public/images/gigachat.svg diff --git a/ui/app/workspace/providers/fragments/apiKeysFormFragment.tsx b/ui/app/workspace/providers/fragments/apiKeysFormFragment.tsx index 244c0fcb99d..bf1105943e8 100644 --- a/ui/app/workspace/providers/fragments/apiKeysFormFragment.tsx +++ b/ui/app/workspace/providers/fragments/apiKeysFormFragment.tsx @@ -15,7 +15,7 @@ import { Control, UseFormReturn } from "react-hook-form"; import { DeploymentsTable } from "./deploymentsTable"; // Providers that support batch APIs -const BATCH_SUPPORTED_PROVIDERS = ["openai", "bedrock", "anthropic", "gemini", "azure", "vertex", "wafer"]; +const BATCH_SUPPORTED_PROVIDERS = ["openai", "bedrock", "anthropic", "gemini", "azure", "vertex", "gigachat", "wafer"]; interface Props { control: Control; @@ -147,6 +147,7 @@ export function ApiKeyFormFragment({ control, providerName, baseProviderType, fo const isVLLM = effectiveProvider === "vllm"; const isOllama = effectiveProvider === "ollama"; const isSGL = effectiveProvider === "sgl"; + const isGigaChat = effectiveProvider === "gigachat"; const isDeepseek = effectiveProvider === "deepseek"; const isFireworks = effectiveProvider === "fireworks"; const isKeylessProvider = isOllama || isSGL; @@ -164,6 +165,8 @@ export function ApiKeyFormFragment({ control, providerName, baseProviderType, fo // Auth type state for Vertex: 'service_account', 'service_account_json', or 'api_key' const [vertexAuthType, setVertexAuthType] = useState<"service_account" | "service_account_json" | "api_key">("service_account"); + const [gigaChatAuthType, setGigaChatAuthType] = useState<"credentials" | "access_token" | "password" | "mtls">("credentials"); + // Detect auth type from existing form values when editing useEffect(() => { if (form.formState.isDirty) return; @@ -225,6 +228,31 @@ export function ApiKeyFormFragment({ control, providerName, baseProviderType, fo useEffect(() => { if (form.formState.isDirty) return; + if (isGigaChat) { + const config = form.getValues("key.gigachat_key_config"); + const accessToken = config?.access_token; + const apiKey = form.getValues("key.value"); + const user = config?.user; + const password = config?.password; + const certFile = config?.cert_file; + const keyFile = config?.key_file; + let detected: "credentials" | "access_token" | "password" | "mtls" = "credentials"; + if (accessToken?.value || accessToken?.ref || apiKey?.value || apiKey?.ref) { + detected = "access_token"; + if (!accessToken?.value && !accessToken?.ref && apiKey) { + form.setValue("key.gigachat_key_config.access_token", apiKey); + } + } else if (user?.value || user?.ref || password?.value || password?.ref) { + detected = "password"; + } else if (certFile || keyFile) { + detected = "mtls"; + } + setGigaChatAuthType(detected); + form.setValue("key.gigachat_key_config._auth_type", detected); + if (!config?.scope) { + form.setValue("key.gigachat_key_config.scope", "GIGACHAT_API_PERS"); + } + } if (isBedrockMantle) { const accessKey = form.getValues("key.bedrock_mantle_key_config.access_key"); const secretKey = form.getValues("key.bedrock_mantle_key_config.secret_key"); @@ -242,7 +270,7 @@ export function ApiKeyFormFragment({ control, providerName, baseProviderType, fo } // form.formState.defaultValues is a dependency so detection re-runs when ProviderKeyForm // repopulates an existing key via form.reset(...) after mount, not only on first render. - }, [isBedrockMantle, form, form.formState.defaultValues]); + }, [isGigaChat, isBedrockMantle, form, form.formState.defaultValues]); return (
@@ -315,7 +343,7 @@ export function ApiKeyFormFragment({ control, providerName, baseProviderType, fo />
{/* Hide API Key field for providers with dedicated auth tabs */} - {!isAzure && !isBedrock && !isBedrockMantle && !isVertex && ( + {!isAzure && !isBedrock && !isBedrockMantle && !isVertex && !isGigaChat && ( } )} + {isGigaChat && ( +
+ +
+ Authentication Method + + Bifrost manages bearer Authorization when configured. mTLS certificates authenticate API requests during TLS. + + { + const next = v as "credentials" | "access_token" | "password" | "mtls"; + setGigaChatAuthType(next); + form.setValue("key.gigachat_key_config._auth_type", next, { shouldDirty: true, shouldValidate: true }); + form.setValue("key.value", undefined, { shouldDirty: true }); + if (next !== "credentials") { + form.setValue("key.gigachat_key_config.credentials", undefined, { shouldDirty: true }); + } + if (next !== "access_token") { + form.setValue("key.gigachat_key_config.access_token", undefined, { shouldDirty: true }); + } + if (next !== "password") { + form.setValue("key.gigachat_key_config.user", undefined, { shouldDirty: true }); + form.setValue("key.gigachat_key_config.password", undefined, { shouldDirty: true }); + } + }} + > + + + OAuth + + + Token + + + Password + + + mTLS + + + +
+ + {gigaChatAuthType === "credentials" && ( + <> + ( + + OAuth Credentials (Required) + + + + + GigaChat authorization key. Bifrost exchanges it for access tokens and refreshes them automatically. + + + + )} + /> + ( + + Scope + + + + Defaults to GIGACHAT_API_PERS when left empty. + + + )} + /> + + )} + + {gigaChatAuthType === "access_token" && ( + ( + + Access Token (Required) + + + + Short-lived bearer token. Bifrost sends it as-is and does not refresh it. + + + )} + /> + )} + + {gigaChatAuthType === "password" && ( + <> + ( + + User (Required) + + + + + + )} + /> + ( + + Password (Required) + + + + + Requires a GigaChat deployment that exposes the SDK-compatible password token endpoint. + + + + )} + /> + + )} + + +
+ ( + + Base URL (Optional) + + + + Overrides the provider base URL. Password auth posts to this URL's /v1/token endpoint. + + + )} + /> + ( + + Auth URL (Optional) + + + + OAuth token exchange URL. Used only by the OAuth credentials method. + + + )} + /> + ( + + Cert File (Optional) + + + + Client certificate path for mTLS. Must be configured together with Key File. + + + )} + /> + ( + + Key File (Optional) + + + + Client private key path for mTLS. Must be configured together with Cert File. + + + )} + /> + ( + + CA Bundle File (Optional) + + + + CA bundle path for GigaChat TLS verification. This does not authenticate requests. + + + )} + /> +
+
+ )} {isReplicate && (
diff --git a/ui/app/workspace/providers/views/modelProviderKeysTableView.tsx b/ui/app/workspace/providers/views/modelProviderKeysTableView.tsx index 76293bdb724..8d93a9644f3 100644 --- a/ui/app/workspace/providers/views/modelProviderKeysTableView.tsx +++ b/ui/app/workspace/providers/views/modelProviderKeysTableView.tsx @@ -92,8 +92,9 @@ export default function ModelProviderKeysTableView({ provider, className, header const providerName = provider.name?.toLowerCase() ?? ""; const isVLLM = providerName === "vllm"; const isOllamaOrSGL = providerName === "ollama" || providerName === "sgl"; - const entityLabel = isVLLM ? "model" : isOllamaOrSGL ? "server" : "key"; - const entityLabelPlural = isVLLM ? "models" : isOllamaOrSGL ? "servers" : "keys"; + const isGigaChat = providerName === "gigachat"; + const entityLabel = isVLLM ? "model" : isOllamaOrSGL ? "server" : isGigaChat ? "credential" : "key"; + const entityLabelPlural = isVLLM ? "models" : isOllamaOrSGL ? "servers" : isGigaChat ? "credentials" : "keys"; const EntityLabel = entityLabel.charAt(0).toUpperCase() + entityLabel.slice(1); const hasUpdateProviderAccess = useRbac(RbacResource.ModelProvider, RbacOperation.Update); const hasDeleteProviderAccess = useRbac(RbacResource.ModelProvider, RbacOperation.Delete); @@ -256,7 +257,7 @@ export default function ModelProviderKeysTableView({ provider, className, header - {isVLLM ? "Model" : isOllamaOrSGL ? "Server" : "API Key"} + {isVLLM ? "Model" : isOllamaOrSGL ? "Server" : isGigaChat ? "Credential" : "API Key"} Weight Enabled @@ -298,7 +299,7 @@ export default function ModelProviderKeysTableView({ provider, className, header )} {key.status === "list_models_failed" && (() => { - // Check if the failure might be due to an env var that the server couldn't resolve + // Check if the failure might be due to a secret reference that the server couldn't resolve const hasSecretVarConfig = (key.azure_key_config?.endpoint?.type && key.azure_key_config.endpoint.type !== "plain_text") || (key.vertex_key_config?.project_id?.type && key.vertex_key_config.project_id.type !== "plain_text") || @@ -306,11 +307,15 @@ export default function ModelProviderKeysTableView({ provider, className, header (key.bedrock_key_config?.region?.type && key.bedrock_key_config.region.type !== "plain_text") || (key.bedrock_mantle_key_config?.region?.type && key.bedrock_mantle_key_config.region.type !== "plain_text") || (key.vllm_key_config?.url?.type && key.vllm_key_config.url.type !== "plain_text") || + (key.gigachat_key_config?.credentials?.type && key.gigachat_key_config.credentials.type !== "plain_text") || + (key.gigachat_key_config?.access_token?.type && key.gigachat_key_config.access_token.type !== "plain_text") || + (key.gigachat_key_config?.user?.type && key.gigachat_key_config.user.type !== "plain_text") || + (key.gigachat_key_config?.password?.type && key.gigachat_key_config.password.type !== "plain_text") || (key.value?.type && key.value.type !== "plain_text"); - const isEnvResolutionError = + const isSecretResolutionError = hasSecretVarConfig && key.description && /not set|empty|missing/i.test(key.description); - return isEnvResolutionError ? ( + return isSecretResolutionError ? (