Skip to main content

sac/model/
client.rs

1use super::*;
2use std::time::Instant;
3
4#[derive(Clone)]
5pub struct ModelClient {
6    client: Client,
7    base_url: String,
8    api_key: String,
9    pub model: String,
10    backend: BackendKind,
11    reasoning_effort: Option<ReasoningEffort>,
12    reasoning_summary: Option<ReasoningSummary>,
13    reasoning_context: Option<ReasoningContext>,
14    event_sink: Option<EventSink>,
15    thread_name: Option<String>,
16}
17
18impl ModelClient {
19    pub fn from_env() -> Result<Self> {
20        Self::from_env_with_overrides(ClientOverrides::default())
21    }
22
23    pub fn from_env_with_overrides(overrides: ClientOverrides) -> Result<Self> {
24        let requested_backend = overrides.backend.unwrap_or(BackendKind::Auto);
25        let base_url = overrides.base_url.unwrap_or_else(|| {
26            std::env::var("OPENAI_BASE_URL").unwrap_or_else(|_| {
27                default_base_url_for_backend_hint(requested_backend).to_string()
28            })
29        });
30        let backend = match requested_backend {
31            BackendKind::Auto => detect_backend(&base_url)?,
32            explicit => explicit,
33        };
34        let api_key = api_key_for_backend(
35            backend,
36            overrides.api_key_env.as_deref(),
37            overrides.api_key.as_deref(),
38        )?;
39        let model = overrides.model.unwrap_or_else(|| {
40            std::env::var("OPENAI_MODEL").unwrap_or_else(|_| default_model_for_backend(backend))
41        });
42        let reasoning_effort = match backend {
43            BackendKind::DeepSeekChat => None,
44            _ => overrides
45                .reasoning_effort
46                .or_else(|| default_reasoning_effort(backend)),
47        };
48        let reasoning_summary = match backend {
49            BackendKind::DeepSeekChat | BackendKind::FireworksChat => None,
50            _ => overrides.reasoning_summary,
51        };
52        let reasoning_context = match backend {
53            BackendKind::DeepSeekChat | BackendKind::FireworksChat => None,
54            _ => overrides.reasoning_context,
55        };
56
57        tracing::debug!(
58            requested_backend = ?requested_backend,
59            resolved_backend = ?backend,
60            backend_source = if matches!(requested_backend, BackendKind::Auto) {
61                "auto_detect"
62            } else {
63                "explicit"
64            },
65            model = %model,
66            reasoning_effort = ?reasoning_effort,
67            reasoning_summary = ?reasoning_summary,
68            reasoning_context = ?reasoning_context,
69            "resolved model client configuration"
70        );
71
72        Ok(Self {
73            client: Client::new(),
74            base_url,
75            api_key,
76            model,
77            backend,
78            reasoning_effort,
79            reasoning_summary,
80            reasoning_context,
81            event_sink: None,
82            thread_name: None,
83        })
84    }
85
86    pub async fn send_turn(
87        &self,
88        messages: Vec<Message>,
89        tools: Vec<ToolDefinition>,
90    ) -> Result<ModelTurnResponse> {
91        let started = Instant::now();
92        let message_count = messages.len();
93        let tool_count = tools.len();
94        tracing::info!(
95            backend = ?self.backend,
96            model = %self.model,
97            reasoning_effort = ?self.reasoning_effort,
98            message_count,
99            tool_count,
100            "starting model turn"
101        );
102
103        let response = match self.backend {
104            BackendKind::Auto => unreachable!("backend auto should be resolved at client creation"),
105            BackendKind::DeepSeekChat => self.send_deepseek_chat(messages, tools).await,
106            BackendKind::FireworksChat => self.send_fireworks_chat(messages, tools).await,
107            BackendKind::OpenAiResponses => self.send_openai_responses(messages, tools).await,
108            BackendKind::ChatGptCodexResponses => {
109                chatgpt_codex::send_responses(
110                    &self.client,
111                    &self.base_url,
112                    &self.model,
113                    self.reasoning_effort.as_ref(),
114                    self.reasoning_summary.as_ref(),
115                    self.reasoning_context.as_ref(),
116                    messages,
117                    tools,
118                )
119                .await
120            }
121        }?;
122
123        tracing::info!(
124            backend = ?self.backend,
125            model = %self.model,
126            finish_reason = ?response.finish_reason,
127            has_text = response.assistant.content.is_some(),
128            tool_call_count = response
129                .assistant
130                .tool_calls
131                .as_ref()
132                .map(|calls| calls.len())
133                .unwrap_or(0),
134            latency_ms = started.elapsed().as_millis() as u64,
135            "model turn completed"
136        );
137
138        Ok(response)
139    }
140
141    pub async fn complete_text(
142        &self,
143        system_prompt: &str,
144        user_prompt: &str,
145    ) -> Result<TextCompletion> {
146        let messages = vec![
147            Message::System {
148                content: system_prompt.to_string(),
149            },
150            Message::User {
151                content: user_prompt.to_string(),
152            },
153        ];
154
155        let response = self.send_turn(messages, Vec::new()).await?;
156        let content = response
157            .assistant
158            .content
159            .ok_or_else(|| anyhow!("Text completion returned no text content"))?;
160
161        Ok(TextCompletion {
162            content,
163            usage: response.usage,
164        })
165    }
166
167    pub fn base_url(&self) -> &str {
168        &self.base_url
169    }
170
171    pub fn backend(&self) -> BackendKind {
172        self.backend
173    }
174
175    pub fn reasoning_effort(&self) -> Option<&ReasoningEffort> {
176        self.reasoning_effort.as_ref()
177    }
178
179    pub fn set_event_sink(&mut self, sink: EventSink) {
180        self.event_sink = Some(sink);
181    }
182
183    pub fn set_thread_name(&mut self, name: Option<String>) {
184        self.thread_name = name;
185    }
186
187    async fn send_fireworks_chat(
188        &self,
189        messages: Vec<Message>,
190        tools: Vec<ToolDefinition>,
191    ) -> Result<ModelTurnResponse> {
192        let url = format!("{}/chat/completions", self.base_url);
193        let mut request = json!({
194            "model": self.model,
195            "messages": messages
196                .iter()
197                .map(fireworks_message_to_value)
198                .collect::<Vec<_>>(),
199            "tools": tools,
200            "temperature": 0.0
201        });
202
203        if let Some(effort) = &self.reasoning_effort {
204            match effort {
205                ReasoningEffort::Low | ReasoningEffort::Medium | ReasoningEffort::High => {
206                    request["reasoning_effort"] = Value::String(effort.as_str().to_string());
207                }
208                unsupported => {
209                    return Err(anyhow!(
210                        "reasoning effort '{}' is not supported by fireworks-chat; use low, medium, or high",
211                        unsupported.as_str()
212                    ));
213                }
214            }
215        }
216
217        tracing::debug!(
218            backend = ?self.backend,
219            endpoint = "chat_completions",
220            request_bytes = json_value_len_bytes(&request)?,
221            "built fireworks chat request"
222        );
223        let value = self.post_json_with_retry(&url, &request).await?;
224        parse_chat_completions_response(&value, &url)
225    }
226
227    async fn send_deepseek_chat(
228        &self,
229        messages: Vec<Message>,
230        tools: Vec<ToolDefinition>,
231    ) -> Result<ModelTurnResponse> {
232        let url = format!("{}/chat/completions", self.base_url);
233        let request = deepseek_chat_request(&self.model, &messages, &tools);
234
235        tracing::debug!(
236            backend = ?self.backend,
237            endpoint = "chat_completions",
238            request_bytes = json_value_len_bytes(&request)?,
239            "built deepseek chat request"
240        );
241        let value = self.post_json_with_retry(&url, &request).await?;
242        parse_chat_completions_response(&value, &url)
243    }
244
245    async fn send_openai_responses(
246        &self,
247        messages: Vec<Message>,
248        tools: Vec<ToolDefinition>,
249    ) -> Result<ModelTurnResponse> {
250        let url = format!("{}/responses", self.base_url);
251        let mut request = json!({
252            "model": self.model,
253            "input": responses_input_items(&messages),
254        });
255
256        if !tools.is_empty() {
257            request["tools"] = Value::Array(
258                tools
259                    .iter()
260                    .map(openai_responses_tool_to_value)
261                    .collect::<Vec<_>>(),
262            );
263        }
264
265        if let Some(effort) = &self.reasoning_effort {
266            let mut reasoning = json!({
267                "effort": effort.as_str(),
268            });
269            if let Some(summary) = &self.reasoning_summary {
270                reasoning["summary"] = json!(summary.as_str());
271            }
272            if let Some(context) = &self.reasoning_context {
273                reasoning["context"] = json!(context.as_str());
274            }
275            request["reasoning"] = reasoning;
276            request["include"] = json!(["reasoning.encrypted_content"]);
277        }
278
279        // Use streaming path when event_sink is available
280        if self.event_sink.is_some() {
281            request["stream"] = json!(true);
282            tracing::debug!(
283                backend = ?self.backend,
284                endpoint = "responses",
285                request_bytes = json_value_len_bytes(&request)?,
286                streaming = true,
287                "built openai responses request (streaming)"
288            );
289            return self.post_streaming_openai_responses(&url, &request).await;
290        }
291
292        tracing::debug!(
293            backend = ?self.backend,
294            endpoint = "responses",
295            request_bytes = json_value_len_bytes(&request)?,
296            "built openai responses request"
297        );
298        let value = self.post_json_with_retry(&url, &request).await?;
299        parse_openai_responses_response(&value, &url)
300    }
301
302    async fn post_streaming_openai_responses(
303        &self,
304        url: &str,
305        body: &Value,
306    ) -> Result<ModelTurnResponse> {
307        let event_sink = self.event_sink.as_ref().unwrap();
308        let thread_name = self.thread_name.clone();
309        let mut last_error = anyhow!("No attempts made");
310        let request_bytes = json_value_len_bytes(body)?;
311
312        for attempt in 0..3 {
313            if attempt > 0 {
314                let delay_secs = 1u64 << (attempt - 1);
315                tracing::warn!(
316                    backend = ?self.backend,
317                    endpoint = "responses_stream",
318                    attempt = attempt + 1,
319                    backoff_secs = delay_secs,
320                    "retrying streaming model HTTP request after backoff"
321                );
322                sleep(Duration::from_secs(delay_secs)).await;
323            }
324
325            let attempt_started = Instant::now();
326            tracing::debug!(
327                backend = ?self.backend,
328                endpoint = "responses_stream",
329                attempt = attempt + 1,
330                request_bytes,
331                "starting streaming model HTTP attempt"
332            );
333
334            let response = self
335                .client
336                .post(url)
337                .header("Authorization", format!("Bearer {}", self.api_key))
338                .header("Content-Type", "application/json")
339                .json(body)
340                .send()
341                .await
342                .map_err(|e| anyhow!("HTTP request failed for {}: {}", url, e))?;
343
344            let status = response.status();
345
346            if status.is_success() {
347                tracing::info!(
348                    backend = ?self.backend,
349                    endpoint = "responses_stream",
350                    attempt = attempt + 1,
351                    status = status.as_u16(),
352                    request_bytes,
353                    latency_ms = attempt_started.elapsed().as_millis() as u64,
354                    "streaming model HTTP connected"
355                );
356
357                // Read the stream incrementally
358                let result = self
359                    .read_sse_stream(response, event_sink, &thread_name, url)
360                    .await;
361
362                event_sink.emit(AgentEvent::StreamComplete {
363                    thread_name: thread_name.clone(),
364                });
365
366                return result;
367            }
368
369            // Read the error body for non-success status
370            let error_body = response
371                .text()
372                .await
373                .unwrap_or_else(|_| "[failed to read body]".to_string());
374
375            if status.as_u16() == 429 || status.is_server_error() {
376                tracing::warn!(
377                    backend = ?self.backend,
378                    endpoint = "responses_stream",
379                    attempt = attempt + 1,
380                    status = status.as_u16(),
381                    request_bytes,
382                    latency_ms = attempt_started.elapsed().as_millis() as u64,
383                    retryable = true,
384                    "streaming model HTTP attempt failed with retryable status"
385                );
386                last_error = anyhow!(
387                    "HTTP {} from {}: {}",
388                    status.as_u16(),
389                    url,
390                    &error_body[..error_body.len().min(500)]
391                );
392                continue;
393            }
394
395            tracing::error!(
396                backend = ?self.backend,
397                endpoint = "responses_stream",
398                attempt = attempt + 1,
399                status = status.as_u16(),
400                request_bytes,
401                latency_ms = attempt_started.elapsed().as_millis() as u64,
402                retryable = false,
403                "streaming model HTTP attempt failed with non-retryable status"
404            );
405            return Err(anyhow!(
406                "HTTP {} from {}: {}",
407                status.as_u16(),
408                url,
409                &error_body[..error_body.len().min(500)]
410            ));
411        }
412
413        Err(last_error)
414    }
415
416    async fn read_sse_stream(
417        &self,
418        response: reqwest::Response,
419        event_sink: &EventSink,
420        thread_name: &Option<String>,
421        url: &str,
422    ) -> Result<ModelTurnResponse> {
423        use futures_util::StreamExt;
424
425        let mut stream = response.bytes_stream();
426        let mut buffer = String::new();
427        let mut output_items: Vec<(usize, Value)> = Vec::new();
428        let mut final_response: Option<Value> = None;
429
430        while let Some(chunk_result) = stream.next().await {
431            let chunk = chunk_result.map_err(|e| anyhow!("Stream read error: {}", e))?;
432            let chunk_str = String::from_utf8_lossy(&chunk);
433            buffer.push_str(&chunk_str);
434
435            // Process complete SSE events (separated by \n\n)
436            while let Some(boundary) = buffer.find("\n\n") {
437                let event_block = buffer[..boundary].to_string();
438                buffer = buffer[boundary + 2..].to_string();
439
440                // Parse the SSE event block
441                let mut event_type = String::new();
442                let mut data_parts: Vec<String> = Vec::new();
443
444                for line in event_block.lines() {
445                    if let Some(value) = line.strip_prefix("event:") {
446                        event_type = value.trim().to_string();
447                    } else if let Some(value) = line.strip_prefix("data:") {
448                        let trimmed = value.trim_start();
449                        data_parts.push(trimmed.to_string());
450                    }
451                }
452
453                if data_parts.is_empty() {
454                    continue;
455                }
456
457                let data = data_parts.join("\n");
458                if data == "[DONE]" {
459                    break;
460                }
461
462                let event: Value = match serde_json::from_str(&data) {
463                    Ok(v) => v,
464                    Err(_) => continue,
465                };
466
467                let etype = event_type.as_str();
468                let json_type = event
469                    .get("type")
470                    .and_then(Value::as_str)
471                    .unwrap_or("");
472
473                match etype.is_empty().then_some(json_type).unwrap_or(etype) {
474                    "response.output_text.delta" => {
475                        if let Some(delta) = event.get("delta").and_then(Value::as_str) {
476                            event_sink.emit(AgentEvent::StreamTextDelta {
477                                thread_name: thread_name.clone(),
478                                text: Some(delta.to_string()),
479                            });
480                        }
481                    }
482                    "response.output_item.done" => {
483                        if let Some(item) = event.get("item").cloned() {
484                            let output_index = event
485                                .get("output_index")
486                                .and_then(Value::as_u64)
487                                .and_then(|i| usize::try_from(i).ok())
488                                .unwrap_or(output_items.len());
489                            output_items.retain(|(idx, _)| *idx != output_index);
490                            output_items.push((output_index, item));
491                        }
492                    }
493                    "response.completed" | "response.done" | "response.incomplete" => {
494                        if let Some(resp) = event.get("response").and_then(Value::as_object) {
495                            let mut response_value = Value::Object(resp.clone());
496                            // If the terminal event has empty output, use accumulated items
497                            let output_is_empty = response_value
498                                .get("output")
499                                .and_then(Value::as_array)
500                                .map(Vec::is_empty)
501                                .unwrap_or(true);
502                            if output_is_empty && !output_items.is_empty() {
503                                output_items.sort_by_key(|(idx, _)| *idx);
504                                response_value["output"] = Value::Array(
505                                    output_items
506                                        .iter()
507                                        .map(|(_, item)| item.clone())
508                                        .collect(),
509                                );
510                            }
511                            final_response = Some(response_value);
512                        }
513                    }
514                    "error" | "response.failed" => {
515                        let msg = event
516                            .get("error")
517                            .and_then(|e| e.get("message"))
518                            .and_then(Value::as_str)
519                            .or_else(|| event.get("message").and_then(Value::as_str))
520                            .unwrap_or("Unknown streaming error");
521                        return Err(anyhow!("Streaming error from {}: {}", url, msg));
522                    }
523                    _ => {}
524                }
525            }
526
527            if final_response.is_some() {
528                break;
529            }
530        }
531
532        // Parse the final response
533        match final_response {
534            Some(value) => parse_openai_responses_response(&value, url),
535            None => {
536                // Fallback: if no terminal event but we have output_items, build response
537                if !output_items.is_empty() {
538                    output_items.sort_by_key(|(idx, _)| *idx);
539                    let constructed = json!({
540                        "status": "completed",
541                        "output": output_items.iter().map(|(_, item)| item.clone()).collect::<Vec<_>>(),
542                        "usage": {"input_tokens": 0, "output_tokens": 0, "total_tokens": 0}
543                    });
544                    parse_openai_responses_response(&constructed, url)
545                } else {
546                    Err(anyhow!(
547                        "SSE stream from {} ended without a terminal response event",
548                        url
549                    ))
550                }
551            }
552        }
553    }
554
555    async fn post_json_with_retry(&self, url: &str, body: &Value) -> Result<Value> {
556        let mut last_error = anyhow!("No attempts made");
557        let request_bytes = json_value_len_bytes(body)?;
558
559        for attempt in 0..3 {
560            if attempt > 0 {
561                let delay_secs = 1u64 << (attempt - 1);
562                tracing::warn!(
563                    backend = ?self.backend,
564                    endpoint = endpoint_name(url),
565                    attempt = attempt + 1,
566                    backoff_secs = delay_secs,
567                    "retrying model HTTP request after backoff"
568                );
569                sleep(Duration::from_secs(delay_secs)).await;
570            }
571
572            let attempt_started = Instant::now();
573            tracing::debug!(
574                backend = ?self.backend,
575                endpoint = endpoint_name(url),
576                attempt = attempt + 1,
577                request_bytes,
578                "starting model HTTP attempt"
579            );
580
581            let response = self
582                .client
583                .post(url)
584                .header("Authorization", format!("Bearer {}", self.api_key))
585                .header("Content-Type", "application/json")
586                .json(body)
587                .send()
588                .await
589                .map_err(|e| anyhow!("HTTP request failed for {}: {}", url, e))?;
590
591            let status = response.status();
592            let body = response
593                .text()
594                .await
595                .map_err(|e| anyhow!("Failed to read response body: {}", e))?;
596            let response_bytes = body.len();
597
598            if status.is_success() {
599                tracing::info!(
600                    backend = ?self.backend,
601                    endpoint = endpoint_name(url),
602                    attempt = attempt + 1,
603                    status = status.as_u16(),
604                    request_bytes,
605                    response_bytes,
606                    latency_ms = attempt_started.elapsed().as_millis() as u64,
607                    "model HTTP attempt succeeded"
608                );
609                return serde_json::from_str::<Value>(&body).map_err(|e| {
610                    tracing::error!(
611                        backend = ?self.backend,
612                        endpoint = endpoint_name(url),
613                        attempt = attempt + 1,
614                        status = status.as_u16(),
615                        response_bytes,
616                        parse_error = %e,
617                        "model HTTP success body failed JSON parse"
618                    );
619                    anyhow!(
620                        "Failed to parse response from {}: {}\nBody: {}",
621                        url,
622                        e,
623                        &body[..body.len().min(500)]
624                    )
625                });
626            }
627
628            if status.as_u16() == 429 || status.is_server_error() {
629                tracing::warn!(
630                    backend = ?self.backend,
631                    endpoint = endpoint_name(url),
632                    attempt = attempt + 1,
633                    status = status.as_u16(),
634                    request_bytes,
635                    response_bytes,
636                    latency_ms = attempt_started.elapsed().as_millis() as u64,
637                    retryable = true,
638                    "model HTTP attempt failed with retryable status"
639                );
640                last_error = anyhow!(
641                    "HTTP {} from {}: {}",
642                    status.as_u16(),
643                    url,
644                    &body[..body.len().min(500)]
645                );
646                continue;
647            }
648
649            tracing::error!(
650                backend = ?self.backend,
651                endpoint = endpoint_name(url),
652                attempt = attempt + 1,
653                status = status.as_u16(),
654                request_bytes,
655                response_bytes,
656                latency_ms = attempt_started.elapsed().as_millis() as u64,
657                retryable = false,
658                "model HTTP attempt failed with non-retryable status"
659            );
660            return Err(anyhow!(
661                "HTTP {} from {}: {}",
662                status.as_u16(),
663                url,
664                &body[..body.len().min(500)]
665            ));
666        }
667
668        Err(last_error)
669    }
670}
671
672fn json_value_len_bytes(value: &Value) -> Result<usize> {
673    Ok(serde_json::to_vec(value)?.len())
674}
675
676fn endpoint_name(url: &str) -> &'static str {
677    if url.contains("/responses") {
678        "responses"
679    } else if url.contains("/chat/completions") {
680        "chat_completions"
681    } else {
682        "unknown"
683    }
684}
685
686#[cfg(test)]
687impl ModelClient {
688    pub fn new_for_test() -> Self {
689        Self {
690            client: reqwest::Client::new(),
691            base_url: "https://api.openai.com/v1".to_string(),
692            api_key: "test_dummy_key".to_string(),
693            model: "gpt-5.5".to_string(),
694            backend: BackendKind::OpenAiResponses,
695            reasoning_effort: Some(ReasoningEffort::Xhigh),
696            reasoning_summary: None,
697            reasoning_context: None,
698            event_sink: None,
699            thread_name: None,
700        }
701    }
702}