1use crate::client::{
7 CacheControl, GenaiProvider, GenerationHints, LLMClient, LLMResponse, TokenUsage,
8};
9use crate::coordinator::{ConversationMessage, MessageRole};
10use ares_types::types::{AppError, ContentPart as AresPart, Result, ToolCall, ToolDefinition};
11use async_trait::async_trait;
12use futures::StreamExt;
13use genai::adapter::AdapterKind;
14use genai::chat::{
15 Binary, CacheControl as GenaiCache, ChatMessage, ChatOptions, ChatRequest, ChatResponse,
16 ChatResponseFormat, ChatStreamEvent, ContentPart, JsonSpec, MessageContent, MessageOptions,
17 ReasoningEffort, Tool, ToolResponse, Usage,
18};
19use genai::resolver::{AuthData, Endpoint};
20use genai::{Client, ModelIden, ServiceTarget};
21use std::sync::RwLock;
22use std::time::Duration;
23
24const PROVIDER_WEB_SEARCH: &str = "provider_web_search";
25
26pub struct GenaiClient {
28 inner: Client,
29 provider: GenaiProvider,
30 hints: RwLock<GenerationHints>,
31}
32
33impl GenaiClient {
34 pub fn new(provider: GenaiProvider) -> Result<Self> {
37 let http = reqwest::Client::builder()
38 .timeout(Duration::from_secs(300))
39 .build()
40 .map_err(|e| AppError::External(format!("failed to build reqwest client: {e}")))?;
41 let inner = Client::builder().with_reqwest(http).build();
42 Ok(Self {
43 inner,
44 provider,
45 hints: RwLock::new(GenerationHints::default()),
46 })
47 }
48
49 fn snapshot_hints(&self) -> GenerationHints {
50 self.hints.read().map(|g| g.clone()).unwrap_or_default()
51 }
52
53 fn effective_kind(&self) -> AdapterKind {
54 rewrite_openai_kind(self.provider.kind, &self.provider.model)
55 }
56
57 fn service_target(&self) -> ServiceTarget {
58 let kind = self.effective_kind();
59 let model = ModelIden::new(kind, self.provider.model.clone());
60 let auth = match kind {
61 AdapterKind::Ollama => AuthData::None,
62 _ => match &self.provider.api_key {
63 Some(key) if !key.is_empty() => AuthData::from_single(key.clone()),
64 _ => AuthData::None,
65 },
66 };
67 let endpoint = Endpoint::from_owned(self.resolve_endpoint(kind));
68 ServiceTarget {
69 endpoint,
70 auth,
71 model,
72 }
73 }
74
75 fn resolve_endpoint(&self, kind: AdapterKind) -> String {
76 if let Some(url) = self
77 .provider
78 .endpoint
79 .as_deref()
80 .map(str::trim)
81 .filter(|s| !s.is_empty())
82 {
83 return ensure_trailing_slash(url);
84 }
85 default_endpoint(
86 kind,
87 self.provider.region.as_deref(),
88 self.provider.vertex_project.as_deref(),
89 self.provider.vertex_location.as_deref(),
90 self.provider.custom_index,
91 )
92 }
93
94 fn chat_options(&self, hints: &GenerationHints, capture_tools: bool) -> ChatOptions {
95 let mut opts = ChatOptions::default().with_capture_usage(true);
96 if capture_tools {
97 opts = opts.with_capture_tool_calls(true);
98 }
99 let max_tokens = hints.max_tokens.or(self.provider.params.max_tokens);
100 if let Some(max) = max_tokens {
101 opts = opts.with_max_tokens(max);
102 }
103 if let Some(temp) = self.provider.params.temperature {
104 opts = opts.with_temperature(f64::from(temp));
105 }
106 if let Some(top_p) = self.provider.params.top_p {
107 opts = opts.with_top_p(f64::from(top_p));
108 }
109 if hints.json_mode {
110 opts = opts.with_response_format(ChatResponseFormat::JsonMode);
111 }
112 if let Some(grammar) = hints.guided_grammar.as_deref() {
113 if let Ok(schema) = serde_json::from_str::<serde_json::Value>(grammar) {
114 if schema.get("type").is_some() {
115 opts = opts.with_response_format(ChatResponseFormat::JsonSpec(JsonSpec::new(
116 "guided", schema,
117 )));
118 }
119 }
120 }
121 if let Some(effort) = hints
122 .reasoning_effort
123 .as_deref()
124 .and_then(ReasoningEffort::from_keyword)
125 {
126 opts = opts.with_reasoning_effort(effort);
127 }
128 if let Some(key) = hints.prompt_cache_key.as_ref() {
129 opts = opts.with_prompt_cache_key(key.clone());
130 }
131 if let Some(cc) = hints.cache_control {
132 opts = opts.with_cache_control(map_cache(cc));
133 }
134 if !self.provider.headers.is_empty() {
135 opts = opts.with_extra_headers(self.provider.headers.clone());
136 }
137 let mut extra = serde_json::Map::new();
138 if let Some(fp) = self.provider.params.frequency_penalty {
139 extra.insert("frequency_penalty".into(), serde_json::json!(fp));
140 }
141 if let Some(pp) = self.provider.params.presence_penalty {
142 extra.insert("presence_penalty".into(), serde_json::json!(pp));
143 }
144 if hints.suppress_reasoning {
145 extra.insert(
146 "chat_template_kwargs".into(),
147 serde_json::json!({ "enable_thinking": false }),
148 );
149 }
150 if let Some(grammar) = hints.guided_grammar.as_deref() {
151 if serde_json::from_str::<serde_json::Value>(grammar)
152 .ok()
153 .and_then(|v| v.get("type").cloned())
154 .is_none()
155 {
156 extra.insert(
157 "guided_grammar".into(),
158 serde_json::Value::String(grammar.to_string()),
159 );
160 }
161 }
162 if !extra.is_empty() {
163 opts = opts.with_extra_body(serde_json::Value::Object(extra));
164 }
165 opts
166 }
167
168 async fn exec(
169 &self,
170 request: ChatRequest,
171 hints: &GenerationHints,
172 capture_tools: bool,
173 ) -> Result<LLMResponse> {
174 let target = self.service_target();
175 let options = self.chat_options(hints, capture_tools);
176 let response = self
177 .inner
178 .exec_chat(target, request, Some(&options))
179 .await
180 .map_err(map_error)?;
181 Ok(map_response(response))
182 }
183
184 async fn exec_stream(
185 &self,
186 request: ChatRequest,
187 hints: &GenerationHints,
188 ) -> Result<Box<dyn futures::Stream<Item = Result<String>> + Send + Unpin>> {
189 let target = self.service_target();
190 let options = self.chat_options(hints, false);
191 let response = self
192 .inner
193 .exec_chat_stream(target, request, Some(&options))
194 .await
195 .map_err(map_error)?;
196 let mut inner = response.stream;
197 let s = async_stream::stream! {
198 while let Some(ev) = inner.next().await {
199 match ev {
200 Ok(ChatStreamEvent::Chunk(chunk)) => yield Ok(chunk.content),
201 Ok(_) => {}
202 Err(err) => yield Err(map_error(err)),
203 }
204 }
205 };
206 Ok(Box::new(Box::pin(s)))
207 }
208}
209
210#[async_trait]
211impl LLMClient for GenaiClient {
212 async fn generate(&self, prompt: &str) -> Result<String> {
213 Ok(self
214 .generate_with_history(&[("user".into(), prompt.to_string())])
215 .await?
216 .content)
217 }
218
219 async fn generate_with_system(&self, system: &str, prompt: &str) -> Result<String> {
220 Ok(self
221 .generate_with_history(&[
222 ("system".into(), system.to_string()),
223 ("user".into(), prompt.to_string()),
224 ])
225 .await?
226 .content)
227 }
228
229 async fn generate_with_history(&self, messages: &[(String, String)]) -> Result<LLMResponse> {
230 let hints = self.snapshot_hints();
231 let request = request_from_role_content(messages, None, &hints);
232 self.exec(request, &hints, false).await
233 }
234
235 async fn generate_with_tools(
236 &self,
237 prompt: &str,
238 tools: &[ToolDefinition],
239 ) -> Result<LLMResponse> {
240 let hints = self.snapshot_hints();
241 let request =
242 request_from_role_content(&[("user".into(), prompt.to_string())], Some(tools), &hints);
243 self.exec(request, &hints, true).await
244 }
245
246 async fn generate_with_tools_and_history(
247 &self,
248 messages: &[ConversationMessage],
249 tools: &[ToolDefinition],
250 ) -> Result<LLMResponse> {
251 let hints = self.snapshot_hints();
252 let request = request_from_conversation(messages, tools, &hints);
253 self.exec(request, &hints, true).await
254 }
255
256 async fn stream(
257 &self,
258 prompt: &str,
259 ) -> Result<Box<dyn futures::Stream<Item = Result<String>> + Send + Unpin>> {
260 let hints = self.snapshot_hints();
261 let request =
262 request_from_role_content(&[("user".into(), prompt.to_string())], None, &hints);
263 self.exec_stream(request, &hints).await
264 }
265
266 async fn stream_with_system(
267 &self,
268 system: &str,
269 prompt: &str,
270 ) -> Result<Box<dyn futures::Stream<Item = Result<String>> + Send + Unpin>> {
271 let hints = self.snapshot_hints();
272 let request = request_from_role_content(
273 &[
274 ("system".into(), system.to_string()),
275 ("user".into(), prompt.to_string()),
276 ],
277 None,
278 &hints,
279 );
280 self.exec_stream(request, &hints).await
281 }
282
283 async fn stream_with_history(
284 &self,
285 messages: &[(String, String)],
286 ) -> Result<Box<dyn futures::Stream<Item = Result<String>> + Send + Unpin>> {
287 let hints = self.snapshot_hints();
288 let request = request_from_role_content(messages, None, &hints);
289 self.exec_stream(request, &hints).await
290 }
291
292 fn model_name(&self) -> &str {
293 &self.provider.model
294 }
295
296 fn supports_hints(&self) -> bool {
297 true
298 }
299
300 fn set_hints(&self, hints: GenerationHints) {
301 if let Ok(mut slot) = self.hints.write() {
302 *slot = hints;
303 }
304 }
305
306 async fn embed(&self, inputs: &[String]) -> Result<Vec<Vec<f32>>> {
307 if inputs.is_empty() {
308 return Ok(Vec::new());
309 }
310 let target = self.service_target();
311 let response = self
312 .inner
313 .embed_batch(target, inputs.to_vec(), None)
314 .await
315 .map_err(map_error)?;
316 Ok(response.into_vectors())
317 }
318
319 fn supports_vision(&self) -> bool {
320 true
321 }
322
323 fn supports_provider_web_search(&self) -> bool {
324 true
325 }
326}
327
328pub(crate) fn rewrite_openai_kind(kind: AdapterKind, model: &str) -> AdapterKind {
330 if kind != AdapterKind::OpenAI {
331 return kind;
332 }
333 if model.starts_with("gpt-5")
334 || (model.starts_with("gpt") && (model.contains("codex") || model.contains("pro")))
335 {
336 AdapterKind::OpenAIResp
337 } else {
338 kind
339 }
340}
341
342pub(crate) fn join_parts(parts: &[AresPart]) -> String {
344 parts
345 .iter()
346 .filter_map(|part| match part {
347 AresPart::Text { text } => Some(text.as_str()),
348 _ => None,
349 })
350 .collect::<Vec<_>>()
351 .join("")
352}
353
354pub(crate) fn map_tools(tools: &[ToolDefinition], web_search: bool) -> Vec<Tool> {
357 let has_provider = tools.iter().any(|t| t.name == PROVIDER_WEB_SEARCH);
358 let mut out: Vec<Tool> = tools
359 .iter()
360 .filter(|t| t.name != PROVIDER_WEB_SEARCH)
361 .map(|t| {
362 Tool::new(t.name.clone())
363 .with_description(t.description.clone())
364 .with_schema(t.parameters.clone())
365 })
366 .collect();
367 if has_provider || web_search {
368 out.push(Tool::new_web_search());
369 }
370 out
371}
372
373fn request_from_role_content(
374 messages: &[(String, String)],
375 tools: Option<&[ToolDefinition]>,
376 hints: &GenerationHints,
377) -> ChatRequest {
378 let mut system = String::new();
379 let mut chat_messages = Vec::new();
380 for (role, content) in messages {
381 match role.as_str() {
382 "system" => {
383 if !system.is_empty() {
384 system.push('\n');
385 }
386 system.push_str(content);
387 }
388 "assistant" => chat_messages.push(ChatMessage::assistant(content.clone())),
389 "tool" => chat_messages.push(ChatMessage::tool(content.clone())),
390 _ => chat_messages.push(ChatMessage::user(content.clone())),
391 }
392 }
393 finish_request(system, chat_messages, tools, hints, None, None)
394}
395
396fn request_from_conversation(
397 messages: &[ConversationMessage],
398 tools: &[ToolDefinition],
399 hints: &GenerationHints,
400) -> ChatRequest {
401 let mut system = String::new();
402 let mut chat_messages = Vec::new();
403 let mut prev_id = hints.previous_response_id.clone();
404 let mut store = hints.store;
405 for msg in messages {
406 if prev_id.is_none() {
407 prev_id = msg.previous_response_id.clone();
408 }
409 if store.is_none() {
410 store = msg.store;
411 }
412 match msg.role {
413 MessageRole::System => {
414 if !system.is_empty() {
415 system.push('\n');
416 }
417 let text = if msg.parts.is_empty() {
418 msg.content.clone()
419 } else {
420 join_parts(&msg.parts)
421 };
422 system.push_str(&text);
423 }
424 MessageRole::User => {
425 let mut message = ChatMessage::user(parts_to_content(&msg.parts, &msg.content));
426 if let Some(cc) = msg.cache_control {
427 message = message.with_options(MessageOptions {
428 cache_control: Some(map_cache(cc)),
429 });
430 }
431 chat_messages.push(message);
432 }
433 MessageRole::Assistant => {
434 let mut parts = ares_parts_to_genai(&msg.parts, &msg.content);
435 if let Some(reason) = msg.reasoning_content.as_ref() {
436 parts.push(ContentPart::ReasoningContent(reason.clone()));
437 }
438 for call in &msg.tool_calls {
439 parts.push(ContentPart::ToolCall(genai::chat::ToolCall {
440 call_id: call.id.clone(),
441 fn_name: call.name.clone(),
442 fn_arguments: call.arguments.clone(),
443 thought_signatures: None,
444 }));
445 }
446 let mut message = ChatMessage::assistant(MessageContent::from_parts(parts));
447 if let Some(cc) = msg.cache_control {
448 message = message.with_options(MessageOptions {
449 cache_control: Some(map_cache(cc)),
450 });
451 }
452 chat_messages.push(message);
453 }
454 MessageRole::Tool => {
455 let response = ToolResponse::new(
456 msg.tool_call_id.clone().unwrap_or_default(),
457 msg.content.clone(),
458 );
459 chat_messages.push(ChatMessage::from(response));
460 }
461 }
462 }
463 finish_request(system, chat_messages, Some(tools), hints, prev_id, store)
464}
465
466fn finish_request(
467 system: String,
468 messages: Vec<ChatMessage>,
469 tools: Option<&[ToolDefinition]>,
470 hints: &GenerationHints,
471 previous_response_id: Option<String>,
472 store: Option<bool>,
473) -> ChatRequest {
474 let mut request = ChatRequest::new(messages);
475 if !system.is_empty() {
476 request = request.with_system(system);
477 }
478 let mapped = match tools {
479 Some(defs) => map_tools(defs, hints.web_search),
480 None if hints.web_search => vec![Tool::new_web_search()],
481 None => Vec::new(),
482 };
483 if !mapped.is_empty() {
484 request = request.with_tools(mapped);
485 }
486 if let Some(id) = previous_response_id.or_else(|| hints.previous_response_id.clone()) {
487 request = request.with_previous_response_id(id);
488 }
489 if let Some(store) = store.or(hints.store) {
490 request = request.with_store(store);
491 }
492 request
493}
494
495fn parts_to_content(parts: &[AresPart], fallback: &str) -> MessageContent {
496 MessageContent::from_parts(ares_parts_to_genai(parts, fallback))
497}
498
499fn ares_parts_to_genai(parts: &[AresPart], fallback: &str) -> Vec<ContentPart> {
500 if parts.is_empty() {
501 return if fallback.is_empty() {
502 Vec::new()
503 } else {
504 vec![ContentPart::Text(fallback.to_string())]
505 };
506 }
507 parts
508 .iter()
509 .map(|part| match part {
510 AresPart::Text { text } => ContentPart::Text(text.clone()),
511 AresPart::ImageUrl { url } => {
512 ContentPart::Binary(Binary::from_url("image/*", url.clone(), None))
513 }
514 AresPart::ImageBase64 { mime, data } => {
515 ContentPart::Binary(Binary::from_base64(mime.clone(), data.clone(), None))
516 }
517 AresPart::FileUrl { url, mime } => ContentPart::Binary(Binary::from_url(
518 mime.clone()
519 .unwrap_or_else(|| "application/octet-stream".into()),
520 url.clone(),
521 None,
522 )),
523 AresPart::FileBase64 { mime, data, name } => ContentPart::Binary(Binary::from_base64(
524 mime.clone(),
525 data.clone(),
526 name.clone(),
527 )),
528 })
529 .collect()
530}
531
532fn map_response(response: ChatResponse) -> LLMResponse {
533 let tool_calls: Vec<ToolCall> = response
534 .tool_calls()
535 .into_iter()
536 .map(|tc| ToolCall {
537 id: tc.call_id.clone(),
538 name: tc.fn_name.clone(),
539 arguments: tc.fn_arguments.clone(),
540 })
541 .collect();
542 let finish_reason = response
543 .stop_reason
544 .as_ref()
545 .map(|r| r.raw().to_string())
546 .unwrap_or_else(|| {
547 if tool_calls.is_empty() {
548 "stop".into()
549 } else {
550 "tool_calls".into()
551 }
552 });
553 LLMResponse {
554 content: response.first_text().unwrap_or("").to_string(),
555 tool_calls,
556 finish_reason,
557 usage: map_usage(&response.usage),
558 reasoning_content: response.reasoning_content,
559 response_id: response.response_id,
560 }
561}
562
563fn map_usage(usage: &Usage) -> Option<TokenUsage> {
564 if usage.prompt_tokens.is_none()
565 && usage.completion_tokens.is_none()
566 && usage.total_tokens.is_none()
567 {
568 return None;
569 }
570 let prompt = usage.prompt_tokens.unwrap_or(0).max(0) as u32;
571 let completion = usage.completion_tokens.unwrap_or(0).max(0) as u32;
572 let total = usage
573 .total_tokens
574 .map(|n| n.max(0) as u32)
575 .unwrap_or(prompt.saturating_add(completion));
576 Some(TokenUsage {
577 prompt_tokens: prompt,
578 completion_tokens: completion,
579 total_tokens: total,
580 cached_tokens: usage
581 .prompt_tokens_details
582 .as_ref()
583 .and_then(|d| d.cached_tokens)
584 .map(|n| i64::from(n.max(0))),
585 })
586}
587
588fn map_cache(cc: CacheControl) -> GenaiCache {
589 match cc {
590 CacheControl::Ephemeral => GenaiCache::Ephemeral,
591 CacheControl::Ephemeral5m => GenaiCache::Ephemeral5m,
592 CacheControl::Ephemeral24h => GenaiCache::Ephemeral24h,
593 }
594}
595
596fn map_error(err: genai::Error) -> AppError {
597 let status = err.status();
598 let message = match status {
599 Some(code) => format!("HTTP {code}: {err}"),
600 None => err.to_string(),
601 };
602 if status.map(|s| s.as_u16()) == Some(429) {
603 AppError::RateLimited(message)
604 } else {
605 AppError::LLM(message)
606 }
607}
608
609fn ensure_trailing_slash(url: &str) -> String {
610 if url.ends_with('/') {
611 url.to_string()
612 } else {
613 format!("{url}/")
614 }
615}
616
617fn default_endpoint(
618 kind: AdapterKind,
619 region: Option<&str>,
620 vertex_project: Option<&str>,
621 vertex_location: Option<&str>,
622 custom_index: Option<u8>,
623) -> String {
624 match kind {
625 AdapterKind::OpenAI | AdapterKind::OpenAIResp => "https://api.openai.com/v1/".into(),
626 AdapterKind::Gemini => "https://generativelanguage.googleapis.com/v1beta/".into(),
627 AdapterKind::Anthropic => "https://api.anthropic.com/v1/".into(),
628 AdapterKind::MiniMax => "https://api.minimax.io/anthropic/v1/".into(),
629 AdapterKind::Ollama => "http://localhost:11434/".into(),
630 AdapterKind::OllamaCloud => "https://ollama.com/".into(),
631 AdapterKind::Cohere => "https://api.cohere.com/v1/".into(),
632 AdapterKind::Fireworks => "https://api.fireworks.ai/inference/v1/".into(),
633 AdapterKind::Together => "https://api.together.xyz/v1/".into(),
634 AdapterKind::Groq => "https://api.groq.com/openai/v1/".into(),
635 AdapterKind::DeepSeek => "https://api.deepseek.com/v1/".into(),
636 AdapterKind::Xai => "https://api.x.ai/v1/".into(),
637 AdapterKind::Aihubmix => "https://aihubmix.com/v1/".into(),
638 AdapterKind::Kimi => "https://api.moonshot.ai/v1/".into(),
639 AdapterKind::Moonshot => "https://api.moonshot.cn/v1/".into(),
640 AdapterKind::Nebius => "https://api.studio.nebius.ai/v1/".into(),
641 AdapterKind::Mimo => "https://api.mimo.com/openai/v1/".into(),
642 AdapterKind::Zai => "https://api.z.ai/api/paas/v4/".into(),
643 AdapterKind::BigModel => "https://open.bigmodel.cn/api/paas/v4/".into(),
644 AdapterKind::Aliyun => "https://dashscope.aliyuncs.com/compatible-mode/v1/".into(),
645 AdapterKind::QwenCloud => "https://dashscope-intl.aliyuncs.com/compatible-mode/v1/".into(),
646 AdapterKind::OpenRouter => "https://openrouter.ai/api/v1/".into(),
647 AdapterKind::AtlasCloud => "https://api.atlascloud.ai/v1/".into(),
648 AdapterKind::GithubCopilot => "https://models.github.ai/inference/".into(),
649 AdapterKind::OpenCodeGo => "https://opencode.ai/zen/go/v1/".into(),
650 AdapterKind::BedrockApi => {
651 let region = region
652 .map(str::to_string)
653 .or_else(|| std::env::var("AWS_REGION").ok())
654 .or_else(|| std::env::var("AWS_DEFAULT_REGION").ok())
655 .unwrap_or_else(|| "us-east-1".into());
656 format!("https://bedrock-runtime.{region}.amazonaws.com/")
657 }
658 AdapterKind::Vertex => {
659 let project = vertex_project
660 .map(str::to_string)
661 .or_else(|| std::env::var("VERTEX_PROJECT_ID").ok())
662 .unwrap_or_default();
663 match vertex_location
664 .map(str::to_string)
665 .or_else(|| std::env::var("VERTEX_LOCATION").ok())
666 {
667 Some(loc) if !loc.is_empty() && loc != "global" => {
668 format!(
669 "https://{loc}-aiplatform.googleapis.com/v1/projects/{project}/locations/{loc}/"
670 )
671 }
672 _ => format!(
673 "https://aiplatform.googleapis.com/v1/projects/{project}/locations/global/"
674 ),
675 }
676 }
677 AdapterKind::Baidu => "https://qianfan.baidubce.com/v2/".into(),
678 AdapterKind::Omlx => std::env::var("OMLX_ENDPOINT")
679 .ok()
680 .filter(|s| !s.is_empty())
681 .map(|s| ensure_trailing_slash(&s))
682 .unwrap_or_else(|| "http://127.0.0.1:8000/v1/".into()),
683 AdapterKind::Custom(n) => {
684 let idx = custom_index.unwrap_or(n);
685 std::env::var(format!("GENAI_{idx}_ENDPOINT"))
686 .ok()
687 .filter(|s| !s.is_empty())
688 .map(|s| ensure_trailing_slash(&s))
689 .unwrap_or_default()
690 }
691 }
692}
693
694#[cfg(test)]
695mod tests {
696 use super::*;
697
698 #[test]
699 fn gpt5_kind_rewrite_only_for_openai() {
700 assert_eq!(
701 rewrite_openai_kind(AdapterKind::OpenAI, "gpt-5"),
702 AdapterKind::OpenAIResp
703 );
704 assert_eq!(
705 rewrite_openai_kind(AdapterKind::OpenAI, "gpt-5-mini"),
706 AdapterKind::OpenAIResp
707 );
708 assert_eq!(
709 rewrite_openai_kind(AdapterKind::OpenAI, "gpt-4o-codex"),
710 AdapterKind::OpenAIResp
711 );
712 assert_eq!(
713 rewrite_openai_kind(AdapterKind::OpenAI, "gpt-4.1-pro"),
714 AdapterKind::OpenAIResp
715 );
716 assert_eq!(
717 rewrite_openai_kind(AdapterKind::OpenAI, "gpt-4o"),
718 AdapterKind::OpenAI
719 );
720 assert_eq!(
721 rewrite_openai_kind(AdapterKind::Anthropic, "gpt-5"),
722 AdapterKind::Anthropic
723 );
724 assert_eq!(
725 rewrite_openai_kind(AdapterKind::OpenAI, "o3-mini"),
726 AdapterKind::OpenAI
727 );
728 }
729
730 #[test]
731 fn join_parts_concatenates_text() {
732 let parts = vec![
733 AresPart::Text {
734 text: "hello ".into(),
735 },
736 AresPart::ImageUrl {
737 url: "https://example.com/x.png".into(),
738 },
739 AresPart::Text {
740 text: "world".into(),
741 },
742 ];
743 assert_eq!(join_parts(&parts), "hello world");
744 assert_eq!(join_parts(&[]), "");
745 }
746
747 #[test]
748 fn provider_web_search_is_stripped_and_replaced() {
749 let tools = vec![
750 ToolDefinition {
751 name: "lookup".into(),
752 description: "lookup".into(),
753 parameters: serde_json::json!({"type": "object"}),
754 },
755 ToolDefinition {
756 name: PROVIDER_WEB_SEARCH.into(),
757 description: "search".into(),
758 parameters: serde_json::json!({"type": "object"}),
759 },
760 ];
761 let mapped = map_tools(&tools, false);
762 assert_eq!(mapped.len(), 2);
763 assert_eq!(mapped[0].name.as_str(), "lookup");
764 assert!(matches!(mapped[1].name, genai::chat::ToolName::WebSearch));
765 assert!(mapped
766 .iter()
767 .all(|t| t.name.as_str() != PROVIDER_WEB_SEARCH));
768
769 let hint_only = map_tools(&[], true);
770 assert_eq!(hint_only.len(), 1);
771 assert!(matches!(
772 hint_only[0].name,
773 genai::chat::ToolName::WebSearch
774 ));
775
776 let none = map_tools(&[], false);
777 assert!(none.is_empty());
778 }
779}