1use 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
27pub struct GenaiClient {
29 inner: Client,
30 provider: GenaiProvider,
31 hints: RwLock<GenerationHints>,
32}
33
34impl GenaiClient {
35 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
392pub(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
406pub(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
418pub(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 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}