use super::openai_compat::{OpenAiCompatAdapter, OpenAiCompatConfig};
use crate::inference::adapter::InferenceAdapter;
use crate::inference::configurator::ResolvedProvider;
use crate::inference::error::InferenceError;
use crate::inference::registry::ProviderId;
pub const FIREWORKS_BASE_URL: &str = "https://api.fireworks.ai/inference/v1";
pub fn build(
resolved: &ResolvedProvider,
base_url: &str,
) -> Result<Box<dyn InferenceAdapter>, InferenceError> {
let key = resolved
.key()
.ok_or(InferenceError::MissingCredential {
provider: ProviderId::Fireworks,
})?
.clone();
let config = OpenAiCompatConfig {
name: ProviderId::Fireworks.as_str().to_string(),
base_url: base_url.to_string(),
api_key: key,
extra_headers: Vec::new(),
capabilities: *resolved.capabilities(),
};
Ok(Box::new(OpenAiCompatAdapter::new(config)?))
}
pub fn factory(resolved: &ResolvedProvider) -> Result<Box<dyn InferenceAdapter>, InferenceError> {
build(resolved, FIREWORKS_BASE_URL)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::inference::types::{ChatMessage, ChatRequest, SecretString};
const FIREWORKS_MODEL: &str = "accounts/fireworks/models/llama-v3p1-8b-instruct";
fn resolved(key: &str) -> ResolvedProvider {
ResolvedProvider::new(
ProviderId::Fireworks,
"fireworks/llama-v3p1-8b-instruct".to_string(),
Some(SecretString::new(key)),
)
}
#[test]
fn factory_builds_named_adapter() {
let adapter = build(&resolved("fw-test"), FIREWORKS_BASE_URL).expect("built");
assert_eq!(adapter.name(), "fireworks");
assert!(!adapter.wants_detailed_usage());
assert!(!adapter.supports_prompt_caching());
assert!(adapter.supports_native_tools());
}
#[test]
fn missing_key_errors() {
let resolved =
ResolvedProvider::new(ProviderId::Fireworks, "fireworks/llama".to_string(), None);
let Err(err) = build(&resolved, FIREWORKS_BASE_URL) else {
panic!("expected MissingCredential");
};
assert!(matches!(
err,
InferenceError::MissingCredential {
provider: ProviderId::Fireworks
}
));
}
#[tokio::test]
#[ignore = "requires FIREWORKS_API_KEY; skipped in CI"]
async fn live_fireworks_call() {
let Ok(key) = std::env::var("FIREWORKS_API_KEY") else {
eprintln!("FIREWORKS_API_KEY not set — skipping live test");
return;
};
if key.trim().is_empty() {
eprintln!("FIREWORKS_API_KEY is empty — skipping live test");
return;
}
let resolved = ResolvedProvider::new(
ProviderId::Fireworks,
FIREWORKS_MODEL.to_string(),
Some(SecretString::new(key)),
);
let adapter = build(&resolved, FIREWORKS_BASE_URL).expect("build adapter");
let mut req = ChatRequest::new(
FIREWORKS_MODEL,
vec![
ChatMessage::system("You are a concise assistant."),
ChatMessage::user("Reply with exactly the word: pong"),
],
);
req.temperature = Some(0.0);
req.max_tokens = Some(16);
let resp = adapter.chat(&req).await.expect("live chat");
let text = resp.first_text().expect("assistant text");
assert!(!text.is_empty(), "assistant text was empty");
assert!(
resp.usage().prompt_tokens > 0,
"prompt_tokens should be > 0"
);
eprintln!(
"live fireworks ok — text: {text:?}, usage: {:?}",
resp.usage()
);
}
}