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 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 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 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 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 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 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 match final_response {
534 Some(value) => parse_openai_responses_response(&value, url),
535 None => {
536 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}