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