-
Notifications
You must be signed in to change notification settings - Fork 946
Expand file tree
/
Copy pathbatch.go
More file actions
385 lines (331 loc) · 11.3 KB
/
Copy pathbatch.go
File metadata and controls
385 lines (331 loc) · 11.3 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
package gemini
import (
"context"
"fmt"
"net/http"
"strings"
"time"
"github.com/bytedance/sonic"
providerUtils "github.com/maximhq/bifrost/core/providers/utils"
"github.com/maximhq/bifrost/core/schemas"
"github.com/valyala/fasthttp"
)
// ToBifrostBatchStatus converts Gemini batch job state to Bifrost status.
func ToBifrostBatchStatus(geminiState string) schemas.BatchStatus {
switch geminiState {
case GeminiBatchStatePending, GeminiBatchStateRunning:
return schemas.BatchStatusInProgress
case GeminiBatchStateSucceeded:
return schemas.BatchStatusCompleted
case GeminiBatchStateFailed:
return schemas.BatchStatusFailed
case GeminiBatchStateCancelling:
return schemas.BatchStatusCancelling
case GeminiBatchStateCancelled:
return schemas.BatchStatusCancelled
case GeminiBatchStateExpired:
return schemas.BatchStatusExpired
default:
return schemas.BatchStatus(geminiState)
}
}
// parseGeminiTimestamp converts Gemini RFC3339 timestamp to Unix timestamp.
func parseGeminiTimestamp(timestamp string) int64 {
if timestamp == "" {
return 0
}
t, err := time.Parse(time.RFC3339, timestamp)
if err != nil {
return 0
}
return t.Unix()
}
// extractBatchIDFromName extracts the batch ID from the full resource name.
// e.g., "batches/abc123" -> "batches/abc123"
func extractBatchIDFromName(name string) string {
return name
}
// buildBatchRequestItems converts Bifrost batch requests to Gemini format.
func buildBatchRequestItems(requests []schemas.BatchRequestItem) []GeminiBatchRequestItem {
items := make([]GeminiBatchRequestItem, 0, len(requests))
for _, req := range requests {
contents := []Content{}
// Try Body first, then fall back to Params (Anthropic SDK uses Params)
requestData := req.Body
if requestData == nil {
requestData = req.Params
}
// Extract messages from the request data
if requestData != nil {
if msgs, ok := requestData["messages"].([]interface{}); ok {
for _, msg := range msgs {
if msgMap, ok := msg.(map[string]interface{}); ok {
role := "user"
if r, ok := msgMap["role"].(string); ok {
if r == "assistant" {
role = "model"
} else if r == "system" {
// System messages are handled separately in Gemini
continue
} else {
role = r
}
}
parts := []*Part{}
if c, ok := msgMap["content"].(string); ok {
parts = append(parts, &Part{Text: c})
}
contents = append(contents, Content{
Role: role,
Parts: parts,
})
}
}
}
}
item := GeminiBatchRequestItem{
Request: GeminiBatchGenerateContentRequest{
Contents: contents,
},
}
// Add metadata with custom_id as key
if req.CustomID != "" {
item.Metadata = &GeminiBatchMetadata{
Key: req.CustomID,
}
}
items = append(items, item)
}
return items
}
// downloadBatchResultsFile downloads and parses a batch results file from Gemini.
// Returns the parsed result items from the JSONL file and any parse errors encountered.
func (provider *GeminiProvider) downloadBatchResultsFile(ctx context.Context, key schemas.Key, fileName string) ([]schemas.BatchResultItem, []schemas.BatchError, *schemas.BifrostError) {
providerName := provider.GetProviderKey()
// Create request to download the file
req := fasthttp.AcquireRequest()
resp := fasthttp.AcquireResponse()
defer fasthttp.ReleaseRequest(req)
defer fasthttp.ReleaseResponse(resp)
// Build download URL - use the download endpoint with alt=media
// The base URL is like https://generativelanguage.googleapis.com/v1beta
// We need to change it to https://generativelanguage.googleapis.com/download/v1beta
baseURL := strings.Replace(provider.networkConfig.BaseURL, "/v1beta", "/download/v1beta", 1)
// Ensure fileName has proper format
fileID := fileName
if !strings.HasPrefix(fileID, "files/") {
fileID = "files/" + fileID
}
url := fmt.Sprintf("%s/%s:download?alt=media", baseURL, fileID)
provider.logger.Debug("gemini batch results file download url: " + url)
providerUtils.SetExtraHeaders(ctx, req, provider.networkConfig.ExtraHeaders, nil)
req.SetRequestURI(url)
req.Header.SetMethod(http.MethodGet)
if key.Value != "" {
req.Header.Set("x-goog-api-key", key.Value)
}
// Make request
_, bifrostErr := providerUtils.MakeRequestWithContext(ctx, provider.client, req, resp)
if bifrostErr != nil {
return nil, nil, bifrostErr
}
// Handle error response
if resp.StatusCode() != fasthttp.StatusOK {
return nil, nil, parseGeminiError(resp, &providerUtils.RequestMetadata{
Provider: providerName,
RequestType: schemas.BatchResultsRequest,
})
}
body, err := providerUtils.CheckAndDecodeBody(resp)
if err != nil {
return nil, nil, providerUtils.NewBifrostOperationError(schemas.ErrProviderResponseDecode, err, providerName)
}
// Parse JSONL content - each line is a separate JSON object
// Use streaming parser to avoid string conversion and collect parse errors
results := make([]schemas.BatchResultItem, 0)
parseResult := providerUtils.ParseJSONL(body, func(line []byte) error {
var resultLine GeminiBatchFileResultLine
if err := sonic.Unmarshal(line, &resultLine); err != nil {
provider.logger.Warn("gemini batch results file parse error: " + err.Error())
return err
}
customID := resultLine.Key
if customID == "" {
customID = fmt.Sprintf("request-%d", len(results))
}
resultItem := schemas.BatchResultItem{
CustomID: customID,
}
if resultLine.Error != nil {
resultItem.Error = &schemas.BatchResultError{
Code: fmt.Sprintf("%d", resultLine.Error.Code),
Message: resultLine.Error.Message,
}
} else if resultLine.Response != nil {
// Convert the response to a map for the Body field
respBody := make(map[string]interface{})
if len(resultLine.Response.Candidates) > 0 {
candidate := resultLine.Response.Candidates[0]
if candidate.Content != nil && len(candidate.Content.Parts) > 0 {
var textParts []string
for _, part := range candidate.Content.Parts {
if part.Text != "" {
textParts = append(textParts, part.Text)
}
}
if len(textParts) > 0 {
respBody["text"] = strings.Join(textParts, "")
}
}
respBody["finish_reason"] = string(candidate.FinishReason)
}
if resultLine.Response.UsageMetadata != nil {
respBody["usage"] = map[string]interface{}{
"prompt_tokens": resultLine.Response.UsageMetadata.PromptTokenCount,
"completion_tokens": resultLine.Response.UsageMetadata.CandidatesTokenCount,
"total_tokens": resultLine.Response.UsageMetadata.TotalTokenCount,
}
}
resultItem.Response = &schemas.BatchResultResponse{
StatusCode: 200,
Body: respBody,
}
}
results = append(results, resultItem)
return nil
})
return results, parseResult.Errors, nil
}
// extractGeminiUsageMetadata extracts usage metadata (as ints) from Gemini response
func extractGeminiUsageMetadata(geminiResponse *GenerateContentResponse) (int, int, int) {
var inputTokens, outputTokens, totalTokens int
if geminiResponse.UsageMetadata != nil {
usageMetadata := geminiResponse.UsageMetadata
inputTokens = int(usageMetadata.PromptTokenCount)
outputTokens = int(usageMetadata.CandidatesTokenCount)
totalTokens = int(usageMetadata.TotalTokenCount)
}
return inputTokens, outputTokens, totalTokens
}
// ==================== SDK RESPONSE CONVERTERS ====================
// These functions convert Bifrost batch responses to Google GenAI SDK format.
// ToGeminiJobState converts Bifrost batch status to Gemini SDK job state.
func ToGeminiJobState(status schemas.BatchStatus) string {
switch status {
case schemas.BatchStatusValidating:
return GeminiJobStatePending
case schemas.BatchStatusInProgress:
return GeminiJobStateRunning
case schemas.BatchStatusFinalizing:
return GeminiJobStateRunning
case schemas.BatchStatusCompleted:
return GeminiJobStateSucceeded
case schemas.BatchStatusFailed:
return GeminiJobStateFailed
case schemas.BatchStatusCancelling:
return GeminiJobStateCancelling
case schemas.BatchStatusCancelled:
return GeminiJobStateCancelled
case schemas.BatchStatusExpired:
return GeminiJobStateFailed
default:
return GeminiJobStatePending
}
}
// ToGeminiBatchJobResponse converts a BifrostBatchCreateResponse to Gemini SDK format.
func ToGeminiBatchJobResponse(resp *schemas.BifrostBatchCreateResponse) *GeminiBatchJobResponseSDK {
if resp == nil {
return nil
}
result := &GeminiBatchJobResponseSDK{
Name: resp.ID,
State: ToGeminiJobState(resp.Status),
}
// Add metadata if available
if resp.CreatedAt > 0 {
result.Metadata = &GeminiBatchMetadata{
Name: resp.ID,
State: ToGeminiJobState(resp.Status),
CreateTime: time.Unix(resp.CreatedAt, 0).Format(time.RFC3339),
BatchStats: &GeminiBatchStats{
RequestCount: resp.RequestCounts.Total,
PendingRequestCount: resp.RequestCounts.Total - resp.RequestCounts.Completed,
SuccessfulRequestCount: resp.RequestCounts.Completed - resp.RequestCounts.Failed,
},
}
}
return result
}
// ToGeminiBatchRetrieveResponse converts a BifrostBatchRetrieveResponse to Gemini SDK format.
func ToGeminiBatchRetrieveResponse(resp *schemas.BifrostBatchRetrieveResponse) *GeminiBatchJobResponseSDK {
if resp == nil {
return nil
}
result := &GeminiBatchJobResponseSDK{
Name: resp.ID,
State: ToGeminiJobState(resp.Status),
}
// Add metadata
result.Metadata = &GeminiBatchMetadata{
Name: resp.ID,
State: ToGeminiJobState(resp.Status),
CreateTime: time.Unix(resp.CreatedAt, 0).Format(time.RFC3339),
BatchStats: &GeminiBatchStats{
RequestCount: resp.RequestCounts.Total,
PendingRequestCount: resp.RequestCounts.Total - resp.RequestCounts.Completed,
SuccessfulRequestCount: resp.RequestCounts.Completed - resp.RequestCounts.Failed,
},
}
if resp.CompletedAt != nil {
result.Metadata.EndTime = time.Unix(*resp.CompletedAt, 0).Format(time.RFC3339)
}
// Add output file info if available
if resp.OutputFileID != nil {
result.Dest = &GeminiBatchDest{
FileName: *resp.OutputFileID,
}
}
return result
}
// ToGeminiBatchListResponse converts a BifrostBatchListResponse to Gemini SDK format.
func ToGeminiBatchListResponse(resp *schemas.BifrostBatchListResponse) *GeminiBatchListResponseSDK {
if resp == nil {
return nil
}
jobs := make([]GeminiBatchJobResponseSDK, 0, len(resp.Data))
for _, batch := range resp.Data {
job := GeminiBatchJobResponseSDK{
Name: batch.ID,
State: ToGeminiJobState(batch.Status),
}
// Add metadata
job.Metadata = &GeminiBatchMetadata{
Name: batch.ID,
State: ToGeminiJobState(batch.Status),
CreateTime: time.Unix(batch.CreatedAt, 0).Format(time.RFC3339),
BatchStats: &GeminiBatchStats{
RequestCount: batch.RequestCounts.Total,
PendingRequestCount: batch.RequestCounts.Total - batch.RequestCounts.Completed,
SuccessfulRequestCount: batch.RequestCounts.Completed - batch.RequestCounts.Failed,
},
}
jobs = append(jobs, job)
}
result := &GeminiBatchListResponseSDK{
BatchJobs: jobs,
}
if resp.NextCursor != nil {
result.NextPageToken = *resp.NextCursor
}
return result
}
// ToGeminiBatchCancelResponse converts a BifrostBatchCancelResponse to Gemini SDK format.
func ToGeminiBatchCancelResponse(resp *schemas.BifrostBatchCancelResponse) *GeminiBatchJobResponseSDK {
if resp == nil {
return nil
}
return &GeminiBatchJobResponseSDK{
Name: resp.ID,
State: ToGeminiJobState(resp.Status),
}
}