Skip to main content

ares_llm/
genai_client.rs

1//! genai-backed [`LLMClient`](crate::client::LLMClient).
2//!
3//! Every chat/embed call is dispatched with an explicit [`ServiceTarget`]
4//! (never a bare `&str` model name, which genai would silently treat as Ollama).
5
6use 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
26/// HTTP LLM client over `genai` 0.7.
27pub struct GenaiClient {
28    inner: Client,
29    provider: GenaiProvider,
30    hints: RwLock<GenerationHints>,
31}
32
33impl GenaiClient {
34    /// Build a client from a resolved provider. Uses our reqwest 0.13 rustls
35    /// client (300s timeout) rather than genai's default builder (which `.expect`s).
36    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
328/// `type = openai` + gpt-5* / gpt*codex / gpt*pro → OpenAI Responses. Other kinds stay put.
329pub(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
342/// Concatenate text parts; non-text parts are ignored.
343pub(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
354/// Map ARES tool definitions to genai tools, stripping `provider_web_search`
355/// and injecting [`ToolName::WebSearch`] when requested.
356pub(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}