返回 CodeWhale
chat_completions.rs
根目录 / crates / app-server / src / chat_completions.rs
1 //! Provider-neutral `/v1/chat/completions` pass-through endpoint.
2 //!
3 //! This module resolves a model inside the configured provider's authority and
4 //! forwards an OpenAI-compatible request body upstream. Model text is metadata:
5 //! it never selects a provider configuration or credential slot. It does
6 //! **not** import or call any DeepSeek-named client APIs — routing stays in
7 //! neutral config/provider types.
8 //!
9 //! Only providers whose [`WireFormat`] is [`WireFormat::ChatCompletions`] are
10 //! served. Streaming requests are explicitly rejected for now.
11
12 use std::collections::BTreeMap;
13
14 use axum::Json;
15 use axum::extract::State;
16 use axum::http::{HeaderName, StatusCode};
17 use axum::response::IntoResponse;
18 use codewhale_agent::ModelRegistry;
19 use codewhale_config::{
20 ConfigApiKeyValueKind, ConfigToml, ProviderKind, apply_openrouter_vendor,
21 auth_mode_disables_api_key, classify_config_api_key_value, is_upstream_auth_header,
22 provider::WireFormat,
23 provider_base_url_is_official, provider_preserves_custom_base_url_model,
24 route::{LogicalModelRef, RouteError, RouteRequest, RouteResolver},
25 validate_openrouter_vendor,
26 };
27 use serde_json::Value;
28
29 use super::AppState;
30
31 // ── Upstream deadlines ─────────────────────────────────────────────────
32
33 /// Connect budget for the upstream forward. Matches the 10s connect bound
34 /// used by the TUI client's non-streaming requests (vision, the shared
35 /// retry client).
36 const UPSTREAM_CONNECT_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(10);
37
38 /// Total budget for one upstream forward, connect through body end. The
39 /// handler rejects streaming (`stream: true`) and reads the full upstream
40 /// body, so without a client-level total a provider that accepts the
41 /// connection and stalls — or trickles the body — wedges this handler (and
42 /// the caller's connection) indefinitely. 1800s mirrors the TUI client's
43 /// non-streaming envelope for the same request class.
44 const UPSTREAM_TOTAL_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(1800);
45
46 // ── Resolved endpoint ──────────────────────────────────────────────────
47
48 /// Everything needed to forward a single chat-completions request upstream.
49 #[derive(Debug, Clone)]
50 struct ResolvedModelEndpoint {
51 provider: ProviderKind,
52 base_url: String,
53 model: String,
54 api_key: Option<String>,
55 auth_disabled: bool,
56 http_headers: BTreeMap<String, String>,
57 path_suffix: Option<String>,
58 insecure_skip_tls_verify: bool,
59 wire_format: WireFormat,
60 }
61
62 // ── Resolution ─────────────────────────────────────────────────────────
63
64 /// Resolve a provider endpoint from the app configuration + an optional
65 /// `model` field pulled out of the incoming request body.
66 fn resolve_endpoint(
67 config: &ConfigToml,
68 registry: &ModelRegistry,
69 request_model: Option<&str>,
70 ) -> Result<ResolvedModelEndpoint, RouteError> {
71 // The configured provider is route authority. A request's model field can
72 // select only within that provider; it may never switch endpoints or
73 // credential slots by resembling another provider's catalog row.
74 let provider_kind = config.provider;
75 let provider_cfg = config.providers.for_provider(provider_kind);
76 let provider_meta = provider_kind.provider();
77
78 // Base URL: configured → default
79 let base_url = provider_base_url(config, provider_kind);
80 let endpoint_owns_models = endpoint_preserves_raw_model_ids(provider_kind, &base_url);
81
82 // ModelRegistry canonicalizes provider-owned aliases and detects clearly
83 // foreign rows. RouteResolver remains authoritative for the provider-scoped
84 // wire model and custom-endpoint passthrough contract.
85 let raw_selected_model = request_model
86 .filter(|m| !m.trim().is_empty())
87 .map(str::to_string)
88 .or_else(|| provider_cfg.model.clone())
89 .or_else(|| {
90 (provider_kind == ProviderKind::Deepseek)
91 .then(|| config.default_text_model.clone())
92 .flatten()
93 })
94 .unwrap_or_else(|| provider_meta.default_model().to_string());
95 let selected_model = if endpoint_owns_models {
96 raw_selected_model
97 } else {
98 match registry.resolve(Some(&raw_selected_model), Some(provider_kind)) {
99 Ok(resolved)
100 if !resolved.used_fallback && resolved.resolved.provider == provider_kind =>
101 {
102 resolved.resolved.id
103 }
104 Err(_) if registry.is_known_for_other_provider(&raw_selected_model, provider_kind) => {
105 return Err(RouteError::ForeignModelForDirectProvider {
106 provider: provider_kind.as_str().into(),
107 model: raw_selected_model,
108 });
109 }
110 // Registry metadata is advisory. An unknown future id stays in the
111 // selected provider's scope, where RouteResolver either accepts the
112 // provider's pass-through contract or rejects the model. It never
113 // borrows another provider's default or credentials.
114 Ok(_) | Err(_) => raw_selected_model,
115 }
116 };
117 let route = RouteResolver::new().resolve(&RouteRequest {
118 explicit_provider: Some(provider_kind),
119 model_selector: Some(LogicalModelRef::from(selected_model.as_str())),
120 saved_provider_model: None,
121 base_url_override: Some(base_url.clone()),
122 limit_overrides: Vec::new(),
123 })?;
124 let model = route.wire_model_id().as_str().to_string();
125
126 let auth_mode = provider_cfg.auth_mode.as_deref().or_else(|| {
127 (provider_kind == config.provider)
128 .then_some(config.auth_mode.as_deref())
129 .flatten()
130 });
131 let auth_disabled = auth_mode_disables_api_key(auth_mode);
132
133 let configured_api_key = provider_cfg.api_key.as_deref();
134
135 // Provider auth comes only from the resolved endpoint configuration. The
136 // HTTP request's Authorization header authenticates the caller to the local
137 // app-server and is never a provider credential.
138 let api_key = resolve_upstream_api_key(
139 configured_api_key,
140 auth_disabled,
141 provider_base_url_is_official(provider_kind, &base_url),
142 || {
143 provider_meta
144 .env_vars()
145 .iter()
146 .find_map(|var| std::env::var(var).ok())
147 },
148 );
149
150 let mut http_headers = if provider_kind == config.provider {
151 config.http_headers.clone()
152 } else {
153 BTreeMap::new()
154 };
155 http_headers.extend(provider_cfg.http_headers.clone());
156 if auth_disabled {
157 http_headers.retain(|name, _| !is_upstream_auth_header(name));
158 }
159
160 let path_suffix = provider_cfg.path_suffix.clone();
161
162 let insecure_skip_tls_verify = provider_cfg.insecure_skip_tls_verify.unwrap_or(false);
163
164 let wire_format = route.protocol();
165
166 Ok(ResolvedModelEndpoint {
167 provider: provider_kind,
168 base_url,
169 model,
170 api_key,
171 auth_disabled,
172 http_headers,
173 path_suffix,
174 insecure_skip_tls_verify,
175 wire_format,
176 })
177 }
178
179 fn resolve_upstream_api_key(
180 configured: Option<&str>,
181 auth_disabled: bool,
182 allow_ambient: bool,
183 ambient_provider_env: impl FnOnce() -> Option<String>,
184 ) -> Option<String> {
185 if auth_disabled {
186 None
187 } else if let Some(configured) = configured
188 .filter(|value| classify_config_api_key_value(value) == ConfigApiKeyValueKind::Literal)
189 {
190 Some(configured.to_string())
191 } else if allow_ambient {
192 ambient_provider_env()
193 } else {
194 None
195 }
196 }
197
198 fn provider_base_url(config: &ConfigToml, provider: ProviderKind) -> String {
199 let metadata = provider.provider();
200 config
201 .providers
202 .for_provider(provider)
203 .base_url
204 .clone()
205 .unwrap_or_else(|| metadata.default_base_url().to_string())
206 }
207
208 fn endpoint_preserves_raw_model_ids(provider: ProviderKind, base_url: &str) -> bool {
209 matches!(
210 provider,
211 ProviderKind::Custom
212 | ProviderKind::Ollama
213 | ProviderKind::OllamaCloud
214 | ProviderKind::Vllm
215 | ProviderKind::Sglang
216 | ProviderKind::OpencodeZen
217 ) || provider_preserves_custom_base_url_model(provider, base_url)
218 }
219
220 /// Build the upstream URL. DeepSeek strict function calls are a beta feature,
221 /// so only requests that actually carry `function.strict = true` preserve the
222 /// configured `/beta` route. Ordinary requests continue to use `/v1`.
223 fn upstream_url(endpoint: &ResolvedModelEndpoint, body: &Value) -> String {
224 let base = endpoint.base_url.trim_end_matches('/');
225 match endpoint.path_suffix.as_deref() {
226 Some(suffix) if !suffix.trim().is_empty() => format!(
227 "{}/{}",
228 unversioned_base_url(base),
229 suffix.trim_start_matches('/')
230 ),
231 _ => {
232 let mut versioned = versioned_base_url(base);
233 let deepseek_strict_beta = endpoint.provider == ProviderKind::Deepseek
234 && provider_base_url_is_official(endpoint.provider, base)
235 && versioned
236 .rsplit('/')
237 .next()
238 .is_some_and(|segment| segment.eq_ignore_ascii_case("beta"))
239 && body_uses_strict_tools(body);
240 if !deepseek_strict_beta
241 && versioned
242 .rsplit('/')
243 .next()
244 .is_some_and(|segment| segment.eq_ignore_ascii_case("beta"))
245 {
246 versioned = format!("{}/v1", unversioned_base_url(base));
247 }
248 format!("{}/chat/completions", versioned.trim_end_matches('/'))
249 }
250 }
251 }
252
253 fn body_uses_strict_tools(body: &Value) -> bool {
254 body.get("tools")
255 .and_then(Value::as_array)
256 .is_some_and(|tools| {
257 tools
258 .iter()
259 .any(|tool| tool.pointer("/function/strict").and_then(Value::as_bool) == Some(true))
260 })
261 }
262
263 fn versioned_base_url(base_url: &str) -> String {
264 let trimmed = base_url.trim_end_matches('/');
265 if base_url_has_version_suffix(trimmed) {
266 trimmed.to_string()
267 } else {
268 format!("{trimmed}/v1")
269 }
270 }
271
272 fn unversioned_base_url(base_url: &str) -> String {
273 let trimmed = base_url.trim_end_matches('/');
274 trimmed
275 .rsplit_once('/')
276 .filter(|(_, segment)| is_version_segment(segment))
277 .map(|(base, _)| base)
278 .unwrap_or(trimmed)
279 .to_string()
280 }
281
282 fn base_url_has_version_suffix(trimmed: &str) -> bool {
283 trimmed.rsplit('/').next().is_some_and(is_version_segment)
284 }
285
286 fn is_version_segment(segment: &str) -> bool {
287 segment.eq_ignore_ascii_case("beta")
288 || segment
289 .strip_prefix('v')
290 .or_else(|| segment.strip_prefix('V'))
291 .is_some_and(|rest| !rest.is_empty() && rest.chars().all(|ch| ch.is_ascii_digit()))
292 }
293
294 // ── Route handler ──────────────────────────────────────────────────────
295
296 pub(crate) async fn chat_completions_handler(
297 State(state): State<AppState>,
298 Json(mut body): Json<Value>,
299 ) -> impl IntoResponse {
300 // Reject streaming early.
301 if body
302 .get("stream")
303 .and_then(|v| v.as_bool())
304 .unwrap_or(false)
305 {
306 return (
307 StatusCode::BAD_REQUEST,
308 Json(serde_json::json!({
309 "error": {
310 "message": "streaming is not supported on this endpoint",
311 "type": "unsupported_parameter",
312 "code": "streaming_unsupported"
313 }
314 })),
315 )
316 .into_response();
317 }
318
319 // Extract model from body.
320 let request_model = body.get("model").and_then(|v| v.as_str());
321
322 // Resolve endpoint. Everything the upstream call needs is copied out of
323 // the config here and the read guard is released before any network
324 // I/O: a slow upstream must never hold the shared config lock, because
325 // a queued `app/config/set` writer would then stall every later reader.
326 let config = state.config.read().await;
327 let vendor = config
328 .providers
329 .for_provider(config.provider)
330 .vendor
331 .as_deref()
332 .unwrap_or_default();
333 let openrouter_vendor = match validate_openrouter_vendor(vendor) {
334 Ok(vendor) if vendor.is_none() || config.provider == ProviderKind::Openrouter => {
335 vendor.map(str::to_owned)
336 }
337 _ => {
338 return (
339 StatusCode::BAD_REQUEST,
340 Json(serde_json::json!({
341 "error": {
342 "message": "vendor is supported only for OpenRouter and must be a slug without whitespace or control characters",
343 "type": "invalid_request_error",
344 "code": "invalid_vendor"
345 }
346 })),
347 )
348 .into_response();
349 }
350 };
351 let resolved = resolve_endpoint(&config, &state.registry, request_model);
352 drop(config);
353 let endpoint = match resolved {
354 Ok(endpoint) => endpoint,
355 Err(error) => {
356 return (
357 StatusCode::BAD_REQUEST,
358 Json(serde_json::json!({
359 "error": {
360 "message": format!("model route could not be resolved: {error}"),
361 "type": "invalid_request_error",
362 "code": "model_route_invalid"
363 }
364 })),
365 )
366 .into_response();
367 }
368 };
369
370 // Only ChatCompletions providers are supported.
371 if endpoint.wire_format != WireFormat::ChatCompletions {
372 return (
373 StatusCode::BAD_REQUEST,
374 Json(serde_json::json!({
375 "error": {
376 "message": format!(
377 "provider {:?} uses {:?} wire format, only ChatCompletions is supported",
378 endpoint.provider, endpoint.wire_format
379 ),
380 "type": "unsupported_provider",
381 "code": "provider_wire_format_unsupported"
382 }
383 })),
384 )
385 .into_response();
386 }
387
388 // Always write the resolved model back. Unknown provider-owned ids remain
389 // byte-for-byte passthrough values, while known aliases become their exact
390 // provider wire ids before forwarding.
391 body["model"] = serde_json::Value::String(endpoint.model.clone());
392 // The operator pin overrides caller ordering/fallback preferences while
393 // retaining caller restrictions such as only, ignore, and privacy policy.
394 apply_openrouter_vendor(&mut body, openrouter_vendor.as_deref());
395
396 let url = upstream_url(&endpoint, &body);
397
398 if endpoint.insecure_skip_tls_verify {
399 return (
400 StatusCode::BAD_REQUEST,
401 Json(serde_json::json!({
402 "error": {
403 "message": format!(
404 "TLS certificate verification cannot be disabled for provider {:?}; use SSL_CERT_FILE with a trusted custom CA bundle",
405 endpoint.provider
406 ),
407 "type": "invalid_request_error",
408 "code": "tls_verification_required"
409 }
410 })),
411 )
412 .into_response();
413 }
414
415 // Build upstream request. The shared platform builder sets no timeouts,
416 // so the proxy would hang forever on an accept-and-stall upstream;
417 // bound both the connect and the whole non-streaming round trip.
418 let upstream_req = codewhale_release::platform_http_client_builder()
419 .connect_timeout(UPSTREAM_CONNECT_TIMEOUT)
420 .timeout(UPSTREAM_TOTAL_TIMEOUT)
421 .build()
422 .map_err(|e| {
423 (
424 StatusCode::INTERNAL_SERVER_ERROR,
425 Json(serde_json::json!({
426 "error": {
427 "message": format!("failed to build upstream client: {e}"),
428 "type": "internal_error"
429 }
430 })),
431 )
432 .into_response()
433 })
434 .map(|client| {
435 let mut req = client.post(&url).json(&body);
436
437 if !endpoint.auth_disabled
438 && let Some(key) = endpoint.api_key.as_deref()
439 {
440 req = req.bearer_auth(key);
441 }
442
443 // Forward configured provider headers.
444 for (name, value) in &endpoint.http_headers {
445 if endpoint.auth_disabled && is_upstream_auth_header(name) {
446 continue;
447 }
448 if let Ok(header_name) = HeaderName::from_bytes(name.as_bytes()) {
449 req = req.header(header_name, value.as_str());
450 }
451 }
452
453 req
454 });
455
456 let client = match upstream_req {
457 Ok(client) => client,
458 Err(resp) => return resp,
459 };
460
461 // Execute upstream request.
462 match client.send().await {
463 Ok(upstream_resp) => {
464 let status = upstream_resp.status();
465 let headers = upstream_resp.headers().clone();
466 match upstream_resp.text().await {
467 Ok(body_text) => {
468 let mut response =
469 axum::response::Response::new(axum::body::Body::from(body_text));
470 *response.status_mut() = status;
471 // Forward relevant upstream headers.
472 if let Some(ct) = headers.get("content-type") {
473 response.headers_mut().insert("content-type", ct.clone());
474 }
475 response
476 }
477 Err(e) => (
478 StatusCode::BAD_GATEWAY,
479 Json(serde_json::json!({
480 "error": {
481 "message": format!("failed to read upstream response: {e}"),
482 "type": "upstream_error"
483 }
484 })),
485 )
486 .into_response(),
487 }
488 }
489 Err(e) => (
490 StatusCode::BAD_GATEWAY,
491 Json(serde_json::json!({
492 "error": {
493 "message": format!("upstream request failed: {e}"),
494 "type": "upstream_error"
495 }
496 })),
497 )
498 .into_response(),
499 }
500 }
501
502 // ── Tests ──────────────────────────────────────────────────────────────
503
504 #[cfg(test)]
505 mod tests {
506 use super::*;
507 use axum::body::Body;
508 use axum::http::{Method, Request};
509 use codewhale_config::provider::WireFormat;
510 use std::fs;
511 use tokio::sync::mpsc;
512 use tower::ServiceExt;
513
514 use super::super::{app_router, build_state};
515
516 fn install_crypto_provider() {
517 crate::install_test_crypto_provider();
518 }
519
520 // The proxy forwards with a client built per request, and reqwest gives
521 // no accessor for a built client's budgets, so a behavioral test would
522 // need an accept-and-stall upstream and a multi-second (connect) or
523 // half-hour (total) wait. Pin the values instead: if the handler stops
524 // applying them this test cannot see it, but a silent constant change
525 // or an inversion of the connect/total ordering cannot slip through.
526 #[test]
527 fn upstream_deadlines_are_bounded_and_ordered() {
528 assert_eq!(UPSTREAM_CONNECT_TIMEOUT, std::time::Duration::from_secs(10));
529 assert_eq!(UPSTREAM_TOTAL_TIMEOUT, std::time::Duration::from_secs(1800));
530 assert!(
531 UPSTREAM_TOTAL_TIMEOUT > UPSTREAM_CONNECT_TIMEOUT,
532 "the total forward budget must leave room beyond the connect budget"
533 );
534 // The shared platform builder must accept both bounds; the handler
535 // chains them onto this builder.
536 codewhale_release::platform_http_client_builder()
537 .connect_timeout(UPSTREAM_CONNECT_TIMEOUT)
538 .timeout(UPSTREAM_TOTAL_TIMEOUT)
539 .build()
540 .expect("platform builder accepts the proxy deadlines");
541 }
542
543 /// Start a minimal upstream mock server that echoes back what it received.
544 async fn start_mock_upstream() -> (String, tokio::task::JoinHandle<()>) {
545 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
546 let addr = listener.local_addr().unwrap();
547 let base_url = format!("http://{}:{}", addr.ip(), addr.port());
548
549 let handle = tokio::spawn(async move {
550 let app = axum::Router::new()
551 .route("/v1/chat/completions", axum::routing::post(mock_handler));
552 axum::serve(listener, app).await.unwrap();
553 });
554
555 // Give the server a moment to start.
556 tokio::time::sleep(std::time::Duration::from_millis(50)).await;
557
558 (base_url, handle)
559 }
560
561 async fn mock_handler(
562 headers: axum::http::HeaderMap,
563 Json(body): Json<Value>,
564 ) -> impl axum::response::IntoResponse {
565 let auth = headers
566 .get("authorization")
567 .and_then(|v| v.to_str().ok())
568 .unwrap_or("none");
569
570 let response_body = serde_json::json!({
571 "id": "chatcmpl-mock",
572 "object": "chat.completion",
573 "created": 1234567890,
574 "model": body.get("model").and_then(|v| v.as_str()).unwrap_or("unknown"),
575 "choices": [{
576 "index": 0,
577 "message": {
578 "role": "assistant",
579 "content": format!("echo: received {} messages, auth={auth}",
580 body.get("messages").and_then(|m| m.as_array()).map(|a| a.len()).unwrap_or(0))
581 },
582 "finish_reason": "stop"
583 }],
584 "usage": {
585 "prompt_tokens": 10,
586 "completion_tokens": 5,
587 "total_tokens": 15
588 }
589 });
590
591 (StatusCode::OK, Json(response_body))
592 }
593
594 async fn capturing_mock_handler(
595 axum::extract::State(captured): axum::extract::State<
596 mpsc::UnboundedSender<axum::http::HeaderMap>,
597 >,
598 headers: axum::http::HeaderMap,
599 body: Json<Value>,
600 ) -> impl axum::response::IntoResponse {
601 captured
602 .send(headers.clone())
603 .expect("capture upstream headers");
604 mock_handler(headers, body).await
605 }
606
607 async fn start_capturing_mock_upstream() -> (
608 String,
609 mpsc::UnboundedReceiver<axum::http::HeaderMap>,
610 tokio::task::JoinHandle<()>,
611 ) {
612 let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
613 .await
614 .expect("bind capturing upstream");
615 let addr = listener.local_addr().expect("capturing upstream address");
616 let base_url = format!("http://{}:{}", addr.ip(), addr.port());
617 let (captured_tx, captured_rx) = mpsc::unbounded_channel();
618
619 let handle = tokio::spawn(async move {
620 let app = axum::Router::new()
621 .route(
622 "/v1/chat/completions",
623 axum::routing::post(capturing_mock_handler),
624 )
625 .with_state(captured_tx);
626 axum::serve(listener, app)
627 .await
628 .expect("serve capturing upstream");
629 });
630
631 (base_url, captured_rx, handle)
632 }
633
634 fn app_with_mock_upstream(
635 auth_token: Option<&str>,
636 mock_base_url: &str,
637 ) -> (axum::Router, tempfile::TempDir) {
638 app_with_mock_upstream_with_provider_extra(auth_token, mock_base_url, "")
639 }
640
641 fn app_with_mock_upstream_with_provider_extra(
642 auth_token: Option<&str>,
643 mock_base_url: &str,
644 provider_extra: &str,
645 ) -> (axum::Router, tempfile::TempDir) {
646 let tmp = tempfile::tempdir().expect("tempdir");
647 let config_path = tmp.path().join("config.toml");
648 let config_content = format!(
649 r#"
650 provider = "arcee"
651 api_key = "sk-deepseek-secret"
652
653 [providers.arcee]
654 base_url = "{mock_base_url}"
655 model = "trinity-large-thinking"
656 api_key = "arcee-configured-key"
657 {provider_extra}
658 "#
659 );
660 fs::write(&config_path, config_content).expect("write config");
661 let state = build_state(
662 Some(config_path),
663 auth_token.map(std::string::ToString::to_string),
664 )
665 .expect("state");
666 (app_router(state, &[]), tmp)
667 }
668
669 fn app_with_together_mock_upstream(mock_base_url: &str) -> (axum::Router, tempfile::TempDir) {
670 let tmp = tempfile::tempdir().expect("tempdir");
671 let config_path = tmp.path().join("config.toml");
672 let config_content = format!(
673 r#"
674 provider = "together"
675
676 [providers.together]
677 base_url = "{mock_base_url}"
678 api_key = "together-configured-key"
679 "#
680 );
681 fs::write(&config_path, config_content).expect("write config");
682 let state = build_state(Some(config_path), None).expect("state");
683 (app_router(state, &[]), tmp)
684 }
685
686 fn app_with_root_deepseek_mock_upstream(
687 mock_base_url: &str,
688 ) -> (axum::Router, tempfile::TempDir) {
689 let tmp = tempfile::tempdir().expect("tempdir");
690 let config_path = tmp.path().join("config.toml");
691 let config_content = format!(
692 r#"
693 provider = "deepseek"
694 api_key = "root-deepseek-key"
695 base_url = "{mock_base_url}"
696 default_text_model = "root-deepseek-model"
697 http_headers = {{ "X-Root-Route" = "kept" }}
698 "#
699 );
700 fs::write(&config_path, config_content).expect("write config");
701 let state = build_state(Some(config_path), None).expect("state");
702 (app_router(state, &[]), tmp)
703 }
704
705 fn app_with_auth_boundary_mock_upstream(
706 auth_token: &str,
707 mock_base_url: &str,
708 provider_api_key: &str,
709 auth_mode: Option<&str>,
710 include_configured_auth_headers: bool,
711 ) -> (axum::Router, tempfile::TempDir) {
712 let tmp = tempfile::tempdir().expect("tempdir");
713 let config_path = tmp.path().join("config.toml");
714 let auth_mode = auth_mode
715 .map(|mode| format!("auth_mode = {mode:?}"))
716 .unwrap_or_default();
717 let configured_auth_headers = if include_configured_auth_headers {
718 r#"http_headers = { aUtHoRiZaTiOn = "Bearer configured-header-secret", "X-API-Key" = "configured-x-key-secret", "Api-Key" = "configured-key-secret", "Proxy-Authorization" = "Basic configured-proxy-secret", "X-Auth-Token" = "configured-auth-token", "X-Access-Token" = "configured-access-token", "X-Goog-Api-Key" = "configured-google-key", Cookie = "session=secret", "X-Route-Metadata" = "safe" }"#
719 } else {
720 ""
721 };
722 let config_content = format!(
723 r#"
724 provider = "arcee"
725
726 [providers.arcee]
727 base_url = "{mock_base_url}"
728 model = "trinity-large-thinking"
729 api_key = {provider_api_key:?}
730 {auth_mode}
731 {configured_auth_headers}
732 "#
733 );
734 fs::write(&config_path, config_content).expect("write config");
735 let state = build_state(Some(config_path), Some(auth_token.to_string())).expect("state");
736 (app_router(state, &[]), tmp)
737 }
738
739 async fn response_body_json(response: axum::response::Response) -> Value {
740 let bytes = axum::body::to_bytes(response.into_body(), usize::MAX)
741 .await
742 .expect("body bytes");
743 serde_json::from_slice(&bytes).expect("json response")
744 }
745
746 #[tokio::test]
747 async fn openrouter_vendor_forwarding_preserves_pin_and_caller_restrictions() {
748 install_crypto_provider();
749 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
750 let mock_url = format!("http://{}", listener.local_addr().unwrap());
751 let (captured_tx, mut captured_rx) = mpsc::unbounded_channel::<Value>();
752 let upstream = axum::Router::new().route(
753 "/v1/chat/completions",
754 axum::routing::post(move |Json(body): Json<Value>| {
755 let captured = captured_tx.clone();
756 async move {
757 captured.send(body).unwrap();
758 Json(serde_json::json!({"choices": []}))
759 }
760 }),
761 );
762 let upstream_task = tokio::spawn(async move {
763 axum::serve(listener, upstream).await.unwrap();
764 });
765
766 for (provider, vendor, status) in [
767 ("openrouter", "deepinfra/turbo", StatusCode::OK),
768 ("openrouter", "", StatusCode::OK),
769 ("openrouter", "bad vendor fixture", StatusCode::BAD_REQUEST),
770 ("arcee", "deepinfra/turbo", StatusCode::BAD_REQUEST),
771 ("arcee", "", StatusCode::OK),
772 ] {
773 let tmp = tempfile::tempdir().unwrap();
774 let config_path = tmp.path().join("config.toml");
775 let openrouter_vendor = if provider == "openrouter" {
776 vendor
777 } else {
778 "dormant/pin"
779 };
780 let arcee_vendor = if provider == "arcee" { vendor } else { "" };
781 fs::write(&config_path, format!(
782 "provider = {provider:?}\n\
783 [providers.openrouter]\nbase_url = {mock_url:?}\napi_key = \"fixture-openrouter-key\"\nvendor = {openrouter_vendor:?}\n\
784 [providers.arcee]\nbase_url = {mock_url:?}\napi_key = \"fixture-arcee-key\"\nvendor = {arcee_vendor:?}\n"
785 )).unwrap();
786 let state = build_state(Some(config_path), None).unwrap();
787 let app = app_router(state, &[]);
788 let caller_policy = serde_json::json!({
789 "order": ["caller/escape"],
790 "allow_fallbacks": true,
791 "only": ["caller/restriction"],
792 "ignore": ["caller/blocked"],
793 "zdr": true,
794 "data_collection": "deny",
795 "require_parameters": true
796 });
797 let body = serde_json::json!({
798 "model": "fixture/model",
799 "messages": [{"role": "user", "content": "hello"}],
800 "provider": caller_policy
801 });
802 let response = app
803 .oneshot(
804 Request::builder()
805 .method(Method::POST)
806 .uri("/v1/chat/completions")
807 .header("content-type", "application/json")
808 .body(Body::from(serde_json::to_vec(&body).unwrap()))
809 .unwrap(),
810 )
811 .await
812 .unwrap();
813 assert_eq!(response.status(), status, "{provider}: {vendor}");
814 if status == StatusCode::BAD_REQUEST {
815 let error = response_body_json(response).await;
816 assert_eq!(error["error"]["code"], "invalid_vendor");
817 assert!(!error.to_string().contains(vendor));
818 assert!(
819 captured_rx.try_recv().is_err(),
820 "invalid config reached upstream"
821 );
822 } else {
823 let forwarded = captured_rx.try_recv().expect("captured forwarded request");
824 let mut expected = caller_policy;
825 if provider == "openrouter" && !vendor.is_empty() {
826 expected["order"] = serde_json::json!([vendor]);
827 expected["allow_fallbacks"] = serde_json::json!(false);
828 }
829 assert_eq!(forwarded["provider"], expected, "{provider}: {vendor}");
830 assert_eq!(forwarded["model"], "fixture/model");
831 }
832 }
833 upstream_task.abort();
834 }
835
836 #[tokio::test]
837 async fn upstream_request_does_not_hold_the_config_lock() {
838 // A slow upstream used to keep `state.config.read()` alive for the
839 // whole request, so a concurrent `app/config/set` (a writer) blocked
840 // and, with tokio's writer-preferring lock, every later reader
841 // queued behind it.
842 install_crypto_provider();
843 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
844 let mock_url = format!("http://{}", listener.local_addr().unwrap());
845 let (arrived_tx, mut arrived_rx) = mpsc::unbounded_channel::<()>();
846 let release = std::sync::Arc::new(tokio::sync::Notify::new());
847 let upstream_release = release.clone();
848 let upstream = axum::Router::new().route(
849 "/v1/chat/completions",
850 axum::routing::post(move || {
851 let arrived = arrived_tx.clone();
852 let release = upstream_release.clone();
853 async move {
854 arrived.send(()).unwrap();
855 release.notified().await;
856 Json(serde_json::json!({"choices": []}))
857 }
858 }),
859 );
860 let upstream_task = tokio::spawn(async move {
861 axum::serve(listener, upstream).await.unwrap();
862 });
863
864 let tmp = tempfile::tempdir().expect("tempdir");
865 let config_path = tmp.path().join("config.toml");
866 fs::write(
867 &config_path,
868 format!(
869 "provider = \"arcee\"\n[providers.arcee]\nbase_url = {mock_url:?}\nmodel = \"trinity-large-thinking\"\napi_key = \"arcee-configured-key\"\n"
870 ),
871 )
872 .expect("write config");
873 let state = build_state(Some(config_path), None).expect("state");
874 let app = app_router(state.clone(), &[]);
875 let body = serde_json::json!({"messages": [{"role": "user", "content": "hello"}]});
876 let request = tokio::spawn(async move {
877 app.oneshot(
878 Request::builder()
879 .method(Method::POST)
880 .uri("/v1/chat/completions")
881 .header("content-type", "application/json")
882 .body(Body::from(serde_json::to_vec(&body).unwrap()))
883 .unwrap(),
884 )
885 .await
886 .unwrap()
887 });
888
889 tokio::time::timeout(std::time::Duration::from_secs(10), arrived_rx.recv())
890 .await
891 .expect("upstream should receive the request")
892 .expect("arrival signal");
893 assert!(
894 state.config.try_write().is_ok(),
895 "config must be writable while the upstream request is in flight"
896 );
897
898 release.notify_one();
899 let response = tokio::time::timeout(std::time::Duration::from_secs(10), request)
900 .await
901 .expect("request should finish once upstream answers")
902 .expect("join request");
903 assert_eq!(response.status(), StatusCode::OK);
904 upstream_task.abort();
905 }
906
907 #[tokio::test]
908 async fn forwards_messages_and_tools() {
909 install_crypto_provider();
910 let (mock_url, _mock) = start_mock_upstream().await;
911 let (app, _tmp) = app_with_mock_upstream(None, &mock_url);
912
913 let body = serde_json::json!({
914 "model": "trinity-large-thinking",
915 "messages": [
916 {"role": "user", "content": "hello"}
917 ],
918 "tools": [{
919 "type": "function",
920 "function": {
921 "name": "get_weather",
922 "description": "Get weather",
923 "parameters": {"type": "object", "properties": {}}
924 }
925 }],
926 "tool_choice": "auto"
927 });
928
929 let response = app
930 .oneshot(
931 Request::builder()
932 .method(Method::POST)
933 .uri("/v1/chat/completions")
934 .header("content-type", "application/json")
935 .body(Body::from(serde_json::to_vec(&body).unwrap()))
936 .unwrap(),
937 )
938 .await
939 .unwrap();
940
941 assert_eq!(response.status(), StatusCode::OK);
942 let resp_body = response_body_json(response).await;
943 assert_eq!(resp_body["model"], "trinity-large-thinking");
944 assert!(
945 resp_body["choices"][0]["message"]["content"]
946 .as_str()
947 .unwrap()
948 .contains("1 messages")
949 );
950 }
951
952 #[tokio::test]
953 async fn default_model_injected_when_omitted() {
954 install_crypto_provider();
955 let (mock_url, _mock) = start_mock_upstream().await;
956 let (app, _tmp) = app_with_mock_upstream(None, &mock_url);
957
958 let body = serde_json::json!({
959 "messages": [
960 {"role": "user", "content": "hello"}
961 ]
962 });
963
964 let response = app
965 .oneshot(
966 Request::builder()
967 .method(Method::POST)
968 .uri("/v1/chat/completions")
969 .header("content-type", "application/json")
970 .body(Body::from(serde_json::to_vec(&body).unwrap()))
971 .unwrap(),
972 )
973 .await
974 .unwrap();
975
976 assert_eq!(response.status(), StatusCode::OK);
977 let resp_body = response_body_json(response).await;
978 // The mock echoes the model it received; should be the configured default.
979 assert_eq!(resp_body["model"], "trinity-large-thinking");
980 }
981
982 #[tokio::test]
983 async fn root_deepseek_compatibility_fields_reach_the_configured_upstream() {
984 install_crypto_provider();
985 let (mock_url, mut captured, _mock) = start_capturing_mock_upstream().await;
986 let (app, _tmp) = app_with_root_deepseek_mock_upstream(&mock_url);
987
988 let body = serde_json::json!({
989 "messages": [{"role": "user", "content": "hello"}]
990 });
991 let response = app
992 .oneshot(
993 Request::builder()
994 .method(Method::POST)
995 .uri("/v1/chat/completions")
996 .header("content-type", "application/json")
997 .body(Body::from(serde_json::to_vec(&body).unwrap()))
998 .unwrap(),
999 )
1000 .await
1001 .unwrap();
1002
1003 assert_eq!(response.status(), StatusCode::OK);
1004 let response_body = response_body_json(response).await;
1005 assert_eq!(response_body["model"], "root-deepseek-model");
1006 assert!(
1007 response_body["choices"][0]["message"]["content"]
1008 .as_str()
1009 .is_some_and(|content| content.contains("auth=Bearer root-deepseek-key"))
1010 );
1011 let headers = tokio::time::timeout(std::time::Duration::from_secs(1), captured.recv())
1012 .await
1013 .expect("upstream request timeout")
1014 .expect("captured upstream request");
1015 assert_eq!(
1016 headers
1017 .get("x-root-route")
1018 .and_then(|value| value.to_str().ok()),
1019 Some("kept")
1020 );
1021 }
1022
1023 #[tokio::test]
1024 async fn configured_model_preserved_when_provided() {
1025 install_crypto_provider();
1026 let (mock_url, _mock) = start_mock_upstream().await;
1027 let (app, _tmp) = app_with_mock_upstream(None, &mock_url);
1028
1029 let body = serde_json::json!({
1030 "model": "custom-model-v2",
1031 "messages": [
1032 {"role": "user", "content": "hello"}
1033 ]
1034 });
1035
1036 let response = app
1037 .oneshot(
1038 Request::builder()
1039 .method(Method::POST)
1040 .uri("/v1/chat/completions")
1041 .header("content-type", "application/json")
1042 .body(Body::from(serde_json::to_vec(&body).unwrap()))
1043 .unwrap(),
1044 )
1045 .await
1046 .unwrap();
1047
1048 assert_eq!(response.status(), StatusCode::OK);
1049 let resp_body = response_body_json(response).await;
1050 assert_eq!(resp_body["model"], "custom-model-v2");
1051 }
1052
1053 #[test]
1054 fn providerless_together_aliases_are_rejected_under_deepseek_authority() {
1055 let config = ConfigToml::default();
1056 let registry = ModelRegistry::default();
1057
1058 for requested in ["inkling", "together-inkling", "thinkingmachines/inkling"] {
1059 assert!(
1060 matches!(
1061 resolve_endpoint(&config, &registry, Some(requested)),
1062 Err(RouteError::ForeignModelForDirectProvider { .. })
1063 ),
1064 "model text must not switch the configured DeepSeek route to Together: {requested}"
1065 );
1066 }
1067 }
1068
1069 #[test]
1070 fn shared_alias_stays_inside_the_configured_provider() {
1071 let config = ConfigToml {
1072 provider: ProviderKind::Openrouter,
1073 ..ConfigToml::default()
1074 };
1075
1076 let endpoint =
1077 resolve_endpoint(&config, &ModelRegistry::default(), Some("deepseek-v4-pro"))
1078 .expect("configured-provider route");
1079
1080 assert_eq!(endpoint.provider, ProviderKind::Openrouter);
1081 }
1082
1083 #[test]
1084 fn unknown_model_under_zai_never_infers_deepseek_provider_authority() {
1085 let config = ConfigToml {
1086 provider: ProviderKind::Zai,
1087 ..ConfigToml::default()
1088 };
1089 let registry = ModelRegistry::default();
1090
1091 let endpoint = resolve_endpoint(&config, &registry, Some("totally-unknown-model"))
1092 .expect("unknown future id stays inside explicit Z.ai authority");
1093 assert_eq!(endpoint.provider, ProviderKind::Zai);
1094 assert_eq!(endpoint.model, "totally-unknown-model");
1095 assert_eq!(endpoint.api_key, None);
1096 }
1097
1098 #[test]
1099 fn known_deepseek_model_under_zai_is_rejected_instead_of_switching_credentials() {
1100 let mut config = ConfigToml {
1101 provider: ProviderKind::Zai,
1102 ..ConfigToml::default()
1103 };
1104 config.providers.zai.api_key = Some("zai-only-key".to_string());
1105 config.providers.deepseek.api_key = Some("must-not-be-selected".to_string());
1106
1107 assert!(matches!(
1108 resolve_endpoint(
1109 &config,
1110 &ModelRegistry::default(),
1111 Some("deepseek-reasoner")
1112 ),
1113 Err(RouteError::ForeignModelForDirectProvider { .. })
1114 ));
1115 }
1116
1117 #[test]
1118 fn configured_together_authority_canonicalizes_its_own_inkling_aliases() {
1119 let config = ConfigToml {
1120 provider: ProviderKind::Together,
1121 ..ConfigToml::default()
1122 };
1123
1124 for requested in ["inkling", "together-inkling", "thinkingmachines/inkling"] {
1125 let endpoint = resolve_endpoint(&config, &ModelRegistry::default(), Some(requested))
1126 .expect("configured Together route");
1127 assert_eq!(endpoint.provider, ProviderKind::Together, "{requested}");
1128 assert_eq!(endpoint.model, "thinkingmachines/inkling", "{requested}");
1129 }
1130 }
1131
1132 #[test]
1133 fn configured_provider_is_required_for_each_official_alias() {
1134 let registry = ModelRegistry::default();
1135
1136 for (requested, provider, expected) in [
1137 (
1138 "qwen3.7-plus",
1139 ProviderKind::Openrouter,
1140 "qwen/qwen3.7-plus",
1141 ),
1142 ("gpt53-codex", ProviderKind::Openai, "gpt-5.3-codex"),
1143 ("arcee-trinity-mini", ProviderKind::Arcee, "trinity-mini"),
1144 ] {
1145 let config = ConfigToml {
1146 provider,
1147 ..ConfigToml::default()
1148 };
1149 let endpoint = resolve_endpoint(&config, &registry, Some(requested))
1150 .expect("provider-owned known alias route");
1151 assert_eq!(endpoint.provider, provider, "{requested}");
1152 assert_eq!(endpoint.model, expected, "{requested}");
1153 }
1154 }
1155
1156 #[test]
1157 fn default_deepseek_route_cannot_claim_foreign_official_aliases() {
1158 let config = ConfigToml::default();
1159
1160 for requested in ["qwen3.7-plus", "gpt53-codex", "arcee-trinity-mini"] {
1161 assert!(
1162 matches!(
1163 resolve_endpoint(&config, &ModelRegistry::default(), Some(requested)),
1164 Err(RouteError::ForeignModelForDirectProvider { .. })
1165 ),
1166 "{requested} must not select another provider from model text"
1167 );
1168 }
1169 }
1170
1171 #[test]
1172 fn opencode_go_app_route_uses_model_protocol_without_cross_provider_fallback() {
1173 let registry = ModelRegistry::default();
1174 for (model, wire) in [
1175 ("grok-4.5", WireFormat::ChatCompletions),
1176 ("kimi-k3", WireFormat::ChatCompletions),
1177 ("grok-4.6", WireFormat::Responses),
1178 ("gpt-5.6-luna", WireFormat::Responses),
1179 ("minimax-m3", WireFormat::AnthropicMessages),
1180 ("qwen3.8-max", WireFormat::AnthropicMessages),
1181 ] {
1182 for requested in [model.to_string(), format!("opencode-go/{model}")] {
1183 for base_url in [None, Some("https://go-gateway.example.test/v1".into())] {
1184 let mut config = ConfigToml {
1185 provider: ProviderKind::OpencodeGo,
1186 ..ConfigToml::default()
1187 };
1188 config.providers.opencode_go.model = Some(requested.clone());
1189 config.providers.opencode_go.base_url = base_url;
1190 for selection in [None, Some(requested.as_str())] {
1191 let endpoint = resolve_endpoint(&config, &registry, selection)
1192 .expect("documented Go route");
1193 assert_eq!(endpoint.provider, ProviderKind::OpencodeGo);
1194 assert_eq!(endpoint.model, model);
1195 assert_eq!(endpoint.wire_format, wire);
1196 }
1197 }
1198 }
1199 }
1200 for model in ["claude-unproven", "gpt-unlisted", "openai/gpt-5.6-luna"] {
1201 let mut config = ConfigToml {
1202 provider: ProviderKind::OpencodeGo,
1203 ..ConfigToml::default()
1204 };
1205 config.providers.opencode_go.model = Some(model.into());
1206 assert!(resolve_endpoint(&config, &registry, None).is_err());
1207 assert!(resolve_endpoint(&config, &registry, Some(model)).is_err());
1208 config.providers.opencode_go.base_url =
1209 Some("https://go-gateway.example.test/v1".into());
1210 assert!(resolve_endpoint(&config, &registry, Some(model)).is_err());
1211 }
1212 }
1213
1214 #[test]
1215 fn opencode_zen_app_route_uses_the_resolved_model_protocol() {
1216 let config = ConfigToml {
1217 provider: ProviderKind::OpencodeZen,
1218 ..ConfigToml::default()
1219 };
1220 let registry = ModelRegistry::default();
1221
1222 for (model, expected) in [
1223 ("gpt-5.5", WireFormat::Responses),
1224 ("claude-sonnet-4-6", WireFormat::AnthropicMessages),
1225 ("deepseek-v4-pro", WireFormat::ChatCompletions),
1226 ] {
1227 let endpoint = resolve_endpoint(&config, &registry, Some(model))
1228 .unwrap_or_else(|error| panic!("{model} should resolve: {error}"));
1229 assert_eq!(endpoint.provider, ProviderKind::OpencodeZen);
1230 assert_eq!(endpoint.model, model);
1231 assert_eq!(endpoint.wire_format, expected);
1232 }
1233
1234 assert!(matches!(
1235 resolve_endpoint(&config, &registry, Some("gemini-3.1-pro")),
1236 Err(RouteError::UnsupportedModelProtocol { .. })
1237 ));
1238 }
1239
1240 #[test]
1241 fn foreign_model_is_rejected_before_credentials_or_headers_can_cross() {
1242 let mut config = ConfigToml {
1243 provider: ProviderKind::Deepseek,
1244 auth_mode: Some("none".to_string()),
1245 ..ConfigToml::default()
1246 };
1247 config.http_headers.insert(
1248 "X-Root-Route".to_string(),
1249 "must-not-cross-providers".to_string(),
1250 );
1251 config.providers.together.api_key = Some("together-key".to_string());
1252
1253 assert!(matches!(
1254 resolve_endpoint(&config, &ModelRegistry::default(), Some("inkling")),
1255 Err(RouteError::ForeignModelForDirectProvider { .. })
1256 ));
1257 }
1258
1259 #[test]
1260 fn every_official_deepseek_endpoint_canonicalizes_retired_aliases() {
1261 let registry = ModelRegistry::default();
1262 for base_url in [
1263 "https://api.deepseek.com",
1264 "https://api.deepseek.com/v1/",
1265 "https://api.deepseek.com/beta",
1266 ] {
1267 for alias in ["deepseek-chat", "deepseek-reasoner"] {
1268 let mut config = ConfigToml::default();
1269 config.providers.deepseek.base_url = Some(base_url.to_string());
1270 let endpoint = resolve_endpoint(&config, &registry, Some(alias))
1271 .expect("official DeepSeek route");
1272 assert_eq!(endpoint.provider, ProviderKind::Deepseek, "{base_url}");
1273 assert_eq!(endpoint.model, "deepseek-v4-flash", "{base_url} {alias}");
1274 }
1275 }
1276 }
1277
1278 #[test]
1279 fn custom_endpoint_preserves_known_registry_alias_verbatim() {
1280 let mut config = ConfigToml {
1281 provider: ProviderKind::Openrouter,
1282 ..ConfigToml::default()
1283 };
1284 config
1285 .providers
1286 .for_provider_mut(ProviderKind::Openrouter)
1287 .base_url = Some("https://gateway.example.test/v1".to_string());
1288
1289 let endpoint = resolve_endpoint(&config, &ModelRegistry::default(), Some("qwen3.7-plus"))
1290 .expect("custom OpenRouter-compatible route");
1291 assert_eq!(endpoint.provider, ProviderKind::Openrouter);
1292 assert_eq!(endpoint.model, "qwen3.7-plus");
1293 }
1294
1295 #[test]
1296 fn custom_endpoint_never_resolves_ambient_provider_env() {
1297 let ambient_was_read = std::cell::Cell::new(false);
1298 let api_key = resolve_upstream_api_key(None, false, false, || {
1299 ambient_was_read.set(true);
1300 Some("ambient-provider-secret".to_string())
1301 });
1302
1303 assert_eq!(api_key, None);
1304 assert!(!ambient_was_read.get());
1305 for sentinel in [codewhale_config::API_KEYRING_SENTINEL, " __KEYRING__ "] {
1306 assert_eq!(
1307 resolve_upstream_api_key(Some(sentinel), false, false, || unreachable!()),
1308 None
1309 );
1310 assert_eq!(
1311 resolve_upstream_api_key(Some(sentinel), false, true, || Some("ambient".into())),
1312 Some("ambient".to_string())
1313 );
1314 }
1315 }
1316
1317 #[test]
1318 fn disabled_auth_never_resolves_configured_or_ambient_credentials() {
1319 let ambient_was_read = std::cell::Cell::new(false);
1320 let api_key = resolve_upstream_api_key(Some("provider-secret"), true, true, || {
1321 ambient_was_read.set(true);
1322 Some("ambient-provider-secret".to_string())
1323 });
1324
1325 assert_eq!(api_key, None);
1326 assert!(!ambient_was_read.get());
1327 }
1328
1329 #[test]
1330 fn active_custom_endpoint_is_not_hijacked_by_known_foreign_alias() {
1331 let mut config = ConfigToml {
1332 provider: ProviderKind::Arcee,
1333 ..ConfigToml::default()
1334 };
1335 config
1336 .providers
1337 .for_provider_mut(ProviderKind::Arcee)
1338 .base_url = Some("https://gateway.example.test/v1".to_string());
1339
1340 let endpoint = resolve_endpoint(&config, &ModelRegistry::default(), Some("qwen3.7-plus"))
1341 .expect("active custom endpoint route");
1342 assert_eq!(endpoint.provider, ProviderKind::Arcee);
1343 assert_eq!(endpoint.model, "qwen3.7-plus");
1344 }
1345
1346 #[test]
1347 fn official_configured_alias_is_canonicalized_when_model_is_omitted() {
1348 let mut config = ConfigToml {
1349 provider: ProviderKind::Openrouter,
1350 ..ConfigToml::default()
1351 };
1352 config
1353 .providers
1354 .for_provider_mut(ProviderKind::Openrouter)
1355 .model = Some("qwen3.7-plus".to_string());
1356
1357 let endpoint = resolve_endpoint(&config, &ModelRegistry::default(), None)
1358 .expect("configured official alias route");
1359 assert_eq!(endpoint.provider, ProviderKind::Openrouter);
1360 assert_eq!(endpoint.model, "qwen/qwen3.7-plus");
1361 }
1362
1363 #[test]
1364 fn configured_together_inkling_alias_is_normalized_when_model_is_omitted() {
1365 let mut config = ConfigToml {
1366 provider: ProviderKind::Together,
1367 ..ConfigToml::default()
1368 };
1369 config
1370 .providers
1371 .for_provider_mut(ProviderKind::Together)
1372 .model = Some("inkling".to_string());
1373
1374 let endpoint = resolve_endpoint(&config, &ModelRegistry::default(), None)
1375 .expect("configured Inkling route");
1376 assert_eq!(endpoint.provider, ProviderKind::Together);
1377 assert_eq!(endpoint.model, "thinkingmachines/inkling");
1378 }
1379
1380 #[tokio::test]
1381 async fn custom_together_endpoint_preserves_explicit_inkling_model_ids() {
1382 install_crypto_provider();
1383 let (mock_url, _mock) = start_mock_upstream().await;
1384 let (app, _tmp) = app_with_together_mock_upstream(&mock_url);
1385
1386 for requested in ["inkling", "together-inkling", "thinkingmachines/inkling"] {
1387 let body = serde_json::json!({
1388 "model": requested,
1389 "messages": [{"role": "user", "content": "hello"}]
1390 });
1391 let response = app
1392 .clone()
1393 .oneshot(
1394 Request::builder()
1395 .method(Method::POST)
1396 .uri("/v1/chat/completions")
1397 .header("content-type", "application/json")
1398 .body(Body::from(serde_json::to_vec(&body).unwrap()))
1399 .unwrap(),
1400 )
1401 .await
1402 .unwrap();
1403
1404 assert_eq!(response.status(), StatusCode::OK, "{requested}");
1405 let resp_body = response_body_json(response).await;
1406 assert_eq!(resp_body["model"], requested, "{requested}");
1407 }
1408 }
1409
1410 #[tokio::test]
1411 async fn configured_api_key_takes_priority_over_incoming_bearer() {
1412 install_crypto_provider();
1413 let (mock_url, _mock) = start_mock_upstream().await;
1414 let (app, _tmp) = app_with_mock_upstream(None, &mock_url);
1415
1416 let body = serde_json::json!({
1417 "model": "trinity-large-thinking",
1418 "messages": [
1419 {"role": "user", "content": "hello"}
1420 ]
1421 });
1422
1423 // Send with an explicit bearer token, but the configured key should win.
1424 let response = app
1425 .oneshot(
1426 Request::builder()
1427 .method(Method::POST)
1428 .uri("/v1/chat/completions")
1429 .header("content-type", "application/json")
1430 .header("authorization", "Bearer user-provided-secret-key")
1431 .body(Body::from(serde_json::to_vec(&body).unwrap()))
1432 .unwrap(),
1433 )
1434 .await
1435 .unwrap();
1436
1437 assert_eq!(response.status(), StatusCode::OK);
1438 let resp_body = response_body_json(response).await;
1439 let content = resp_body["choices"][0]["message"]["content"]
1440 .as_str()
1441 .unwrap();
1442 // The configured key takes priority, not the incoming Bearer.
1443 assert!(
1444 content.contains("auth=Bearer arcee-configured-key"),
1445 "expected configured auth in mock echo, got: {content}"
1446 );
1447 }
1448
1449 #[tokio::test]
1450 async fn app_authorization_is_not_forwarded_when_upstream_auth_is_disabled() {
1451 install_crypto_provider();
1452 let (mock_url, mut captured, _mock) = start_capturing_mock_upstream().await;
1453 let (app, _tmp) = app_with_auth_boundary_mock_upstream(
1454 "app-secret",
1455 &mock_url,
1456 "provider-secret",
1457 Some("none"),
1458 true,
1459 );
1460
1461 let body = serde_json::json!({
1462 "model": "trinity-large-thinking",
1463 "messages": [{"role": "user", "content": "hello"}]
1464 });
1465 let response = app
1466 .oneshot(
1467 Request::builder()
1468 .method(Method::POST)
1469 .uri("/v1/chat/completions")
1470 .header("content-type", "application/json")
1471 .header("authorization", "Bearer app-secret")
1472 .body(Body::from(serde_json::to_vec(&body).unwrap()))
1473 .unwrap(),
1474 )
1475 .await
1476 .unwrap();
1477
1478 assert_eq!(response.status(), StatusCode::OK);
1479 let headers = tokio::time::timeout(std::time::Duration::from_secs(1), captured.recv())
1480 .await
1481 .expect("upstream request timeout")
1482 .expect("captured upstream request");
1483 for name in [
1484 "authorization",
1485 "x-api-key",
1486 "api-key",
1487 "proxy-authorization",
1488 "x-auth-token",
1489 "x-access-token",
1490 "x-goog-api-key",
1491 "cookie",
1492 ] {
1493 assert!(headers.get(name).is_none(), "disabled auth leaked {name}");
1494 }
1495 assert_eq!(
1496 headers
1497 .get("x-route-metadata")
1498 .and_then(|value| value.to_str().ok()),
1499 Some("safe")
1500 );
1501 }
1502
1503 #[tokio::test]
1504 async fn configured_provider_credential_is_the_only_outbound_bearer() {
1505 install_crypto_provider();
1506 let (mock_url, mut captured, _mock) = start_capturing_mock_upstream().await;
1507 let (app, _tmp) = app_with_auth_boundary_mock_upstream(
1508 "app-secret",
1509 &mock_url,
1510 "provider-secret",
1511 None,
1512 false,
1513 );
1514
1515 let body = serde_json::json!({
1516 "model": "trinity-large-thinking",
1517 "messages": [{"role": "user", "content": "hello"}]
1518 });
1519 let response = app
1520 .oneshot(
1521 Request::builder()
1522 .method(Method::POST)
1523 .uri("/v1/chat/completions")
1524 .header("content-type", "application/json")
1525 .header("authorization", "Bearer app-secret")
1526 .body(Body::from(serde_json::to_vec(&body).unwrap()))
1527 .unwrap(),
1528 )
1529 .await
1530 .unwrap();
1531
1532 assert_eq!(response.status(), StatusCode::OK);
1533 let headers = tokio::time::timeout(std::time::Duration::from_secs(1), captured.recv())
1534 .await
1535 .expect("upstream request timeout")
1536 .expect("captured upstream request");
1537 assert_eq!(
1538 headers
1539 .get("authorization")
1540 .and_then(|value| value.to_str().ok()),
1541 Some("Bearer provider-secret")
1542 );
1543 }
1544
1545 #[tokio::test]
1546 async fn configured_api_key_used_when_no_bearer_in_request() {
1547 install_crypto_provider();
1548 let (mock_url, _mock) = start_mock_upstream().await;
1549 let (app, _tmp) = app_with_mock_upstream(None, &mock_url);
1550
1551 let body = serde_json::json!({
1552 "model": "trinity-large-thinking",
1553 "messages": [
1554 {"role": "user", "content": "hello"}
1555 ]
1556 });
1557
1558 // No Authorization header; the configured key should be used.
1559 let response = app
1560 .oneshot(
1561 Request::builder()
1562 .method(Method::POST)
1563 .uri("/v1/chat/completions")
1564 .header("content-type", "application/json")
1565 .body(Body::from(serde_json::to_vec(&body).unwrap()))
1566 .unwrap(),
1567 )
1568 .await
1569 .unwrap();
1570
1571 assert_eq!(response.status(), StatusCode::OK);
1572 let resp_body = response_body_json(response).await;
1573 let content = resp_body["choices"][0]["message"]["content"]
1574 .as_str()
1575 .unwrap();
1576 assert!(
1577 content.contains("auth=Bearer arcee-configured-key"),
1578 "expected configured auth in mock echo, got: {content}"
1579 );
1580 }
1581
1582 #[tokio::test]
1583 async fn insecure_tls_skip_verify_is_rejected() {
1584 install_crypto_provider();
1585 let (mock_url, _mock) = start_mock_upstream().await;
1586 let (app, _tmp) = app_with_mock_upstream_with_provider_extra(
1587 None,
1588 &mock_url,
1589 "insecure_skip_tls_verify = true",
1590 );
1591
1592 let body = serde_json::json!({
1593 "model": "trinity-large-thinking",
1594 "messages": [
1595 {"role": "user", "content": "hello"}
1596 ]
1597 });
1598
1599 let response = app
1600 .oneshot(
1601 Request::builder()
1602 .method(Method::POST)
1603 .uri("/v1/chat/completions")
1604 .header("content-type", "application/json")
1605 .body(Body::from(serde_json::to_vec(&body).unwrap()))
1606 .unwrap(),
1607 )
1608 .await
1609 .unwrap();
1610
1611 assert_eq!(response.status(), StatusCode::BAD_REQUEST);
1612 let resp_body = response_body_json(response).await;
1613 assert_eq!(resp_body["error"]["code"], "tls_verification_required");
1614 assert!(
1615 resp_body["error"]["message"]
1616 .as_str()
1617 .unwrap()
1618 .contains("SSL_CERT_FILE")
1619 );
1620 }
1621
1622 #[tokio::test]
1623 async fn streaming_request_rejected() {
1624 install_crypto_provider();
1625 let (mock_url, _mock) = start_mock_upstream().await;
1626 let (app, _tmp) = app_with_mock_upstream(None, &mock_url);
1627
1628 let body = serde_json::json!({
1629 "model": "trinity-large-thinking",
1630 "messages": [
1631 {"role": "user", "content": "hello"}
1632 ],
1633 "stream": true
1634 });
1635
1636 let response = app
1637 .oneshot(
1638 Request::builder()
1639 .method(Method::POST)
1640 .uri("/v1/chat/completions")
1641 .header("content-type", "application/json")
1642 .body(Body::from(serde_json::to_vec(&body).unwrap()))
1643 .unwrap(),
1644 )
1645 .await
1646 .unwrap();
1647
1648 assert_eq!(response.status(), StatusCode::BAD_REQUEST);
1649 let resp_body = response_body_json(response).await;
1650 assert_eq!(resp_body["error"]["code"], "streaming_unsupported");
1651 }
1652
1653 #[tokio::test]
1654 async fn requires_bearer_token_when_auth_enabled() {
1655 install_crypto_provider();
1656 let (mock_url, _mock) = start_mock_upstream().await;
1657 let (app, _tmp) = app_with_mock_upstream(Some("test-token"), &mock_url);
1658
1659 let body = serde_json::json!({
1660 "messages": [{"role": "user", "content": "hello"}]
1661 });
1662
1663 let response = app
1664 .oneshot(
1665 Request::builder()
1666 .method(Method::POST)
1667 .uri("/v1/chat/completions")
1668 .header("content-type", "application/json")
1669 .body(Body::from(serde_json::to_vec(&body).unwrap()))
1670 .unwrap(),
1671 )
1672 .await
1673 .unwrap();
1674
1675 assert_eq!(response.status(), StatusCode::UNAUTHORIZED);
1676 }
1677
1678 #[tokio::test]
1679 async fn non_chat_completions_provider_rejected() {
1680 // Use the test to verify WireFormat checks work for non-ChatCompletions providers.
1681 // Anthropic's wire format is AnthropicMessages; OpenaiCodex is Responses.
1682 let endpoint = ResolvedModelEndpoint {
1683 provider: ProviderKind::Anthropic,
1684 base_url: "https://api.anthropic.com".to_string(),
1685 model: "claude-sonnet-4-20250514".to_string(),
1686 api_key: Some("sk-ant-test".to_string()),
1687 auth_disabled: false,
1688 http_headers: BTreeMap::new(),
1689 path_suffix: None,
1690 insecure_skip_tls_verify: false,
1691 wire_format: WireFormat::AnthropicMessages,
1692 };
1693
1694 assert_ne!(endpoint.wire_format, WireFormat::ChatCompletions);
1695 // The handler would reject this; we verify the wire format here.
1696 assert_eq!(endpoint.wire_format, WireFormat::AnthropicMessages);
1697 }
1698
1699 #[test]
1700 fn upstream_url_defaults_to_v1_chat_completions() {
1701 let endpoint = ResolvedModelEndpoint {
1702 provider: ProviderKind::Arcee,
1703 base_url: "https://api.arcee.ai".to_string(),
1704 model: "trinity".to_string(),
1705 api_key: None,
1706 auth_disabled: false,
1707 http_headers: BTreeMap::new(),
1708 path_suffix: None,
1709 insecure_skip_tls_verify: false,
1710 wire_format: WireFormat::ChatCompletions,
1711 };
1712 assert_eq!(
1713 upstream_url(&endpoint, &serde_json::json!({})),
1714 "https://api.arcee.ai/v1/chat/completions"
1715 );
1716 }
1717
1718 #[test]
1719 fn upstream_url_preserves_arcee_api_v1_base() {
1720 let endpoint = ResolvedModelEndpoint {
1721 provider: ProviderKind::Arcee,
1722 base_url: "https://api.arcee.ai/api/v1".to_string(),
1723 model: "trinity".to_string(),
1724 api_key: None,
1725 auth_disabled: false,
1726 http_headers: BTreeMap::new(),
1727 path_suffix: None,
1728 insecure_skip_tls_verify: false,
1729 wire_format: WireFormat::ChatCompletions,
1730 };
1731 assert_eq!(
1732 upstream_url(&endpoint, &serde_json::json!({})),
1733 "https://api.arcee.ai/api/v1/chat/completions"
1734 );
1735 }
1736
1737 #[test]
1738 fn upstream_url_respects_path_suffix() {
1739 let endpoint = ResolvedModelEndpoint {
1740 provider: ProviderKind::Openrouter,
1741 base_url: "https://openrouter.ai/api/v1".to_string(),
1742 model: "deepseek/deepseek-v4-pro".to_string(),
1743 api_key: None,
1744 auth_disabled: false,
1745 http_headers: BTreeMap::new(),
1746 path_suffix: Some("/chat/completions".to_string()),
1747 insecure_skip_tls_verify: false,
1748 wire_format: WireFormat::ChatCompletions,
1749 };
1750 assert_eq!(
1751 upstream_url(&endpoint, &serde_json::json!({})),
1752 "https://openrouter.ai/api/chat/completions"
1753 );
1754 }
1755
1756 #[test]
1757 fn upstream_url_beta_base_uses_v1_for_ordinary_chat_completions() {
1758 let endpoint = ResolvedModelEndpoint {
1759 provider: ProviderKind::Deepseek,
1760 base_url: "https://api.deepseek.com/beta".to_string(),
1761 model: "deepseek-chat".to_string(),
1762 api_key: None,
1763 auth_disabled: false,
1764 http_headers: BTreeMap::new(),
1765 path_suffix: None,
1766 insecure_skip_tls_verify: false,
1767 wire_format: WireFormat::ChatCompletions,
1768 };
1769 assert_eq!(
1770 upstream_url(&endpoint, &serde_json::json!({})),
1771 "https://api.deepseek.com/v1/chat/completions"
1772 );
1773 }
1774
1775 #[test]
1776 fn upstream_url_beta_base_preserves_strict_chat_completions() {
1777 let endpoint = ResolvedModelEndpoint {
1778 provider: ProviderKind::Deepseek,
1779 base_url: "https://api.deepseek.com/beta".to_string(),
1780 model: "deepseek-v4-pro".to_string(),
1781 api_key: None,
1782 auth_disabled: false,
1783 http_headers: BTreeMap::new(),
1784 path_suffix: None,
1785 insecure_skip_tls_verify: false,
1786 wire_format: WireFormat::ChatCompletions,
1787 };
1788 let body = serde_json::json!({
1789 "tools": [{
1790 "type": "function",
1791 "function": {
1792 "name": "lookup",
1793 "strict": true,
1794 "parameters": {"type": "object"}
1795 }
1796 }]
1797 });
1798
1799 assert_eq!(
1800 upstream_url(&endpoint, &body),
1801 "https://api.deepseek.com/beta/chat/completions"
1802 );
1803 }
1804
1805 #[test]
1806 fn upstream_url_strips_trailing_slash() {
1807 let endpoint = ResolvedModelEndpoint {
1808 provider: ProviderKind::Deepseek,
1809 base_url: "https://api.deepseek.com/".to_string(),
1810 model: "deepseek-chat".to_string(),
1811 api_key: None,
1812 auth_disabled: false,
1813 http_headers: BTreeMap::new(),
1814 path_suffix: None,
1815 insecure_skip_tls_verify: false,
1816 wire_format: WireFormat::ChatCompletions,
1817 };
1818 assert_eq!(
1819 upstream_url(&endpoint, &serde_json::json!({})),
1820 "https://api.deepseek.com/v1/chat/completions"
1821 );
1822 }
1823 }
1824
1824 lines RUST