lazyBoy/crates/harness/src/resolve.rs

298 lines
9.7 KiB
Rust
Raw Normal View History

2026-09-04 09:08:56 +00:00
use lazyboy_contracts::{ModelCapabilities, ModelProvider, default_model_id, model_capabilities};
use rig_core::client::CompletionClient;
use rig_core::providers::{openai, xai};
2026-09-03 12:46:14 +00:00
use thiserror::Error;
#[derive(Debug, Error, Clone, PartialEq, Eq)]
pub enum ModelError {
#[error("unsupported_provider:{provider}")]
UnsupportedProvider { provider: String },
#[error("missing credential for {provider} ({env_key})")]
2026-09-04 09:08:56 +00:00
MissingCredential { provider: String, env_key: String },
#[error("missing base URL for {provider}")]
MissingBaseUrl { provider: String },
2026-09-03 12:46:14 +00:00
#[error("unknown model provider: {0}")]
UnknownProvider(String),
#[error("model client error: {0}")]
ProviderClient(String),
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct CredentialChain {
pub bot: Option<String>,
pub space: Option<String>,
pub env: Option<String>,
}
impl CredentialChain {
pub fn resolve(&self) -> Option<&str> {
self.bot
.as_deref()
.filter(|value| !value.is_empty())
.or(self.space.as_deref().filter(|value| !value.is_empty()))
.or(self.env.as_deref().filter(|value| !value.is_empty()))
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ResolveModelRequest {
pub provider: ModelProvider,
pub model_id: Option<String>,
2026-09-04 09:08:56 +00:00
pub base_url: Option<String>,
2026-09-03 12:46:14 +00:00
pub credentials: CredentialChain,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ResolvedBackend {
pub provider: ModelProvider,
pub model_id: String,
pub capabilities: ModelCapabilities,
pub api_key: String,
pub base_url: String,
}
/// Pick a backend without talking to the network.
pub fn resolve_backend(request: ResolveModelRequest) -> Result<ResolvedBackend, ModelError> {
match request.provider {
2026-09-04 09:08:56 +00:00
ModelProvider::Xai | ModelProvider::OpencodeGo | ModelProvider::OpenaiCompatible => {}
2026-09-03 12:46:14 +00:00
other => {
return Err(ModelError::UnsupportedProvider {
provider: other.as_str().to_string(),
});
}
}
2026-09-04 09:08:56 +00:00
let api_key = match request.credentials.resolve() {
Some(key) => key.to_string(),
None if request.provider.requires_api_key() => {
return Err(ModelError::MissingCredential {
provider: request.provider.as_str().to_string(),
env_key: request.provider.env_key_name().to_string(),
});
}
None => String::new(),
};
2026-09-03 12:46:14 +00:00
let model_id = request
.model_id
.filter(|value| !value.is_empty())
.or_else(|| default_model_id(request.provider).map(str::to_string))
2026-09-04 09:08:56 +00:00
.ok_or_else(|| ModelError::ProviderClient("missing model id".into()))?;
2026-09-03 12:46:14 +00:00
let base_url = request
2026-09-04 09:08:56 +00:00
.base_url
.filter(|value| !value.trim().is_empty())
.or_else(|| request.provider.default_base_url().map(str::to_string))
.ok_or_else(|| ModelError::MissingBaseUrl {
provider: request.provider.as_str().to_string(),
})?;
let base_url = rewrite_loopback_host(&normalize_base_url(&base_url));
2026-09-03 12:46:14 +00:00
Ok(ResolvedBackend {
provider: request.provider,
model_id: model_id.clone(),
capabilities: model_capabilities(request.provider, &model_id),
2026-09-04 09:08:56 +00:00
api_key,
2026-09-03 12:46:14 +00:00
base_url,
})
}
2026-09-04 09:08:56 +00:00
fn normalize_base_url(url: &str) -> String {
url.trim().trim_end_matches('/').to_string()
}
fn rewrite_loopback_host(url: &str) -> String {
let supervisor = std::env::var("SANDBOX_SUPERVISOR_URL").unwrap_or_default();
if !supervisor.contains("://supervisor") {
return url.to_string();
}
url.replace("://127.0.0.1", "://host.docker.internal")
.replace("://localhost", "://host.docker.internal")
}
pub enum DynModel {
Xai(xai::completion::CompletionModel),
OpenAi(openai::completion::CompletionModel),
}
pub fn connect_model(backend: &ResolvedBackend) -> Result<DynModel, ModelError> {
match backend.provider {
ModelProvider::Xai => {
let client = xai::Client::new(&backend.api_key)
.map_err(|error| ModelError::ProviderClient(error.to_string()))?;
Ok(DynModel::Xai(client.completion_model(&backend.model_id)))
}
ModelProvider::OpencodeGo | ModelProvider::OpenaiCompatible => {
let key = if backend.api_key.is_empty() {
"local"
} else {
backend.api_key.as_str()
};
let client = openai::CompletionsClient::builder()
.api_key(key.to_string())
.base_url(&backend.base_url)
.build()
.map_err(|error| ModelError::ProviderClient(error.to_string()))?;
Ok(DynModel::OpenAi(client.completion_model(&backend.model_id)))
}
other => Err(ModelError::UnsupportedProvider {
provider: other.as_str().to_string(),
}),
}
}
2026-09-03 12:46:14 +00:00
pub fn credential_from_env(provider: ModelProvider) -> Option<String> {
2026-09-04 09:08:56 +00:00
std::env::var(provider.env_key_name())
.ok()
.filter(|value| !value.is_empty())
2026-09-03 12:46:14 +00:00
}
/// Prove the xAI Rig client can be constructed from a resolved backend.
2026-09-04 09:08:56 +00:00
pub fn connect_xai(
backend: &ResolvedBackend,
) -> Result<rig_core::providers::xai::Client, ModelError> {
2026-09-03 12:46:14 +00:00
if backend.provider != ModelProvider::Xai {
return Err(ModelError::UnsupportedProvider {
provider: backend.provider.as_str().to_string(),
});
}
2026-09-04 09:08:56 +00:00
rig_core::providers::xai::Client::new(&backend.api_key)
.map_err(|error| ModelError::ProviderClient(error.to_string()))
2026-09-03 12:46:14 +00:00
}
#[cfg(test)]
mod tests {
use super::*;
2026-09-04 09:08:56 +00:00
use lazyboy_contracts::DEFAULT_XAI_MODEL;
2026-09-03 12:46:14 +00:00
fn xai_request(key: Option<&str>) -> ResolveModelRequest {
ResolveModelRequest {
provider: ModelProvider::Xai,
model_id: None,
2026-09-04 09:08:56 +00:00
base_url: None,
2026-09-03 12:46:14 +00:00
credentials: CredentialChain {
bot: None,
space: None,
env: key.map(str::to_string),
},
}
}
#[test]
fn xai_resolves_with_default_vision_model() {
let backend = resolve_backend(xai_request(Some("test-key"))).unwrap();
assert_eq!(backend.provider, ModelProvider::Xai);
assert_eq!(backend.model_id, DEFAULT_XAI_MODEL);
assert!(backend.capabilities.vision);
assert_eq!(backend.base_url, "https://api.x.ai/v1");
assert!(connect_xai(&backend).is_ok());
}
#[test]
fn credentials_prefer_bot_then_space_then_env() {
let backend = resolve_backend(ResolveModelRequest {
provider: ModelProvider::Xai,
model_id: Some("grok-4.6".into()),
2026-09-04 09:08:56 +00:00
base_url: None,
2026-09-03 12:46:14 +00:00
credentials: CredentialChain {
bot: Some("bot-key".into()),
space: Some("space-key".into()),
env: Some("env-key".into()),
},
})
.unwrap();
assert_eq!(backend.api_key, "bot-key");
}
#[test]
fn missing_credential_is_explicit() {
let error = resolve_backend(xai_request(None)).unwrap_err();
assert!(matches!(error, ModelError::MissingCredential { .. }));
}
#[test]
fn other_providers_are_reserved_not_silent_fallback() {
for provider in [
ModelProvider::Openai,
ModelProvider::Anthropic,
ModelProvider::Openrouter,
] {
let error = resolve_backend(ResolveModelRequest {
provider,
model_id: Some("whatever".into()),
2026-09-04 09:08:56 +00:00
base_url: None,
2026-09-03 12:46:14 +00:00
credentials: CredentialChain {
bot: None,
space: None,
env: Some("sk-test".into()),
},
})
.unwrap_err();
assert_eq!(
error,
ModelError::UnsupportedProvider {
provider: provider.as_str().to_string()
}
);
}
}
#[test]
fn unknown_provider_strings_fail_to_parse() {
let error = "gemini".parse::<ModelProvider>().unwrap_err();
assert_eq!(error.0, "gemini");
}
2026-09-04 09:08:56 +00:00
#[test]
fn opencode_go_uses_zen_go_endpoint() {
let backend = resolve_backend(ResolveModelRequest {
provider: ModelProvider::OpencodeGo,
model_id: None,
base_url: None,
credentials: CredentialChain {
bot: None,
space: None,
env: Some("go-key".into()),
},
})
.unwrap();
assert_eq!(backend.model_id, "glm-5.1");
assert_eq!(backend.base_url, "https://opencode.ai/zen/go/v1");
assert!(connect_model(&backend).is_ok());
}
#[test]
fn openai_compatible_allows_empty_key_and_custom_url() {
let backend = resolve_backend(ResolveModelRequest {
provider: ModelProvider::OpenaiCompatible,
model_id: Some("qwen2.5".into()),
base_url: Some("http://127.0.0.1:8000/v1/".into()),
credentials: CredentialChain {
bot: None,
space: None,
env: None,
},
})
.unwrap();
assert_eq!(backend.api_key, "");
assert_eq!(backend.base_url, "http://127.0.0.1:8000/v1");
assert!(connect_model(&backend).is_ok());
}
#[test]
fn openai_compatible_requires_base_url() {
let error = resolve_backend(ResolveModelRequest {
provider: ModelProvider::OpenaiCompatible,
model_id: Some("local".into()),
base_url: None,
credentials: CredentialChain {
bot: None,
space: None,
env: None,
},
})
.unwrap_err();
assert!(matches!(error, ModelError::MissingBaseUrl { .. }));
}
2026-09-03 12:46:14 +00:00
}