#![forbid(unsafe_code)]
#![warn(missing_docs)]
pub mod errors;
use serde::{Deserialize, Serialize};
pub use errors::{Result, SdkError};
pub struct Client {
base_url: String,
http: reqwest::Client,
api_key: Option<String>,
}
impl Client {
pub fn new(base_url: &str) -> Result<Self> {
let base_url = base_url.trim_end_matches('/').to_string();
let http = reqwest::Client::builder()
.timeout(std::time::Duration::from_secs(120))
.build()
.map_err(|e| SdkError::Http(e.to_string()))?;
let api_key = std::env::var("ARKNET_WALLET").ok();
Ok(Self {
base_url,
http,
api_key,
})
}
pub async fn connect(opts: ConnectOptions) -> Result<Self> {
let seeds = if opts.seeds.is_empty() {
fetch_seeds().await
} else {
opts.seeds
};
let http = reqwest::Client::builder()
.timeout(std::time::Duration::from_secs(10))
.build()
.map_err(|e| SdkError::Http(e.to_string()))?;
for seed in &seeds {
let url = format!("{}/v1/gateways", seed.trim_end_matches('/'));
let resp = match http.get(&url).send().await {
Ok(r) if r.status().is_success() => r,
_ => continue,
};
let body: serde_json::Value = match resp.json().await {
Ok(v) => v,
Err(_) => continue,
};
let gateways = body["gateways"].as_array().cloned().unwrap_or_default();
let mut sorted = gateways;
sorted.sort_by_key(|g| {
if g["https"].as_bool() == Some(true) {
0
} else {
1
}
});
for gw in &sorted {
let is_https = gw["https"].as_bool() == Some(true);
if opts.require_https && !is_https {
continue;
}
if let Some(gw_url) = gw["url"].as_str() {
return Self::new(gw_url);
}
}
}
Err(SdkError::Http("no reachable gateway found".into()))
}
pub async fn chat_completion(&self, req: ChatRequest) -> Result<ChatResponse> {
let url = format!("{}/v1/chat/completions", self.base_url);
let mut builder = self.http.post(&url).json(&req);
if let Some(key) = &self.api_key {
builder = builder.header("Authorization", format!("Bearer {key}"));
}
let resp = builder
.send()
.await
.map_err(|e| SdkError::Http(e.to_string()))?;
if !resp.status().is_success() {
let status = resp.status().as_u16();
let body = resp.text().await.unwrap_or_default();
return Err(SdkError::Api { status, body });
}
resp.json::<ChatResponse>()
.await
.map_err(|e| SdkError::Http(e.to_string()))
}
pub async fn list_models(&self) -> Result<ModelsResponse> {
let url = format!("{}/v1/models", self.base_url);
let mut builder = self.http.get(&url);
if let Some(key) = &self.api_key {
builder = builder.header("Authorization", format!("Bearer {key}"));
}
let resp = builder
.send()
.await
.map_err(|e| SdkError::Http(e.to_string()))?;
if !resp.status().is_success() {
let status = resp.status().as_u16();
let body = resp.text().await.unwrap_or_default();
return Err(SdkError::Api { status, body });
}
resp.json::<ModelsResponse>()
.await
.map_err(|e| SdkError::Http(e.to_string()))
}
}
const SEEDS_JSON_URL: &str = "https://arknet.arkengel.com/seeds.json";
const FALLBACK_SEEDS: &[&str] = &["https://api.arknet.arkengel.com"];
async fn fetch_seeds() -> Vec<String> {
let client = match reqwest::Client::builder()
.timeout(std::time::Duration::from_secs(5))
.build()
{
Ok(c) => c,
Err(_) => return FALLBACK_SEEDS.iter().map(|s| s.to_string()).collect(),
};
let resp = match client.get(SEEDS_JSON_URL).send().await {
Ok(r) if r.status().is_success() => r,
_ => return FALLBACK_SEEDS.iter().map(|s| s.to_string()).collect(),
};
let body: serde_json::Value = match resp.json().await {
Ok(v) => v,
Err(_) => return FALLBACK_SEEDS.iter().map(|s| s.to_string()).collect(),
};
let urls: Vec<String> = body["seeds"]
.as_array()
.map(|arr| {
arr.iter()
.filter_map(|s| s["url"].as_str().map(String::from))
.collect()
})
.unwrap_or_default();
if urls.is_empty() {
FALLBACK_SEEDS.iter().map(|s| s.to_string()).collect()
} else {
urls
}
}
#[derive(Clone, Debug, Default)]
pub struct ConnectOptions {
pub seeds: Vec<String>,
pub require_https: bool,
}
#[derive(Clone, Debug, Default, Serialize)]
pub struct ChatRequest {
pub model: String,
pub messages: Vec<Message>,
#[serde(skip_serializing_if = "Option::is_none")]
pub max_tokens: Option<u32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub temperature: Option<f64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub stream: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
pub stop: Option<Vec<String>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub prefer_tee: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
pub require_https: Option<bool>,
}
#[derive(Clone, Debug, Default, Serialize, Deserialize)]
pub struct Message {
pub role: String,
pub content: String,
}
#[derive(Clone, Debug, Deserialize)]
pub struct ChatResponse {
pub id: String,
pub choices: Vec<ChatChoice>,
pub usage: Option<TokenUsage>,
}
#[derive(Clone, Debug, Deserialize)]
pub struct ChatChoice {
pub index: u32,
pub message: Message,
pub finish_reason: Option<String>,
}
#[derive(Clone, Debug, Deserialize)]
pub struct TokenUsage {
pub prompt_tokens: u32,
pub completion_tokens: u32,
pub total_tokens: u32,
}
#[derive(Clone, Debug, Deserialize)]
pub struct ModelsResponse {
pub data: Vec<ModelInfo>,
}
#[derive(Clone, Debug, Deserialize)]
pub struct ModelInfo {
pub id: String,
pub owned_by: String,
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn client_trims_trailing_slash() {
let c = Client::new("http://localhost:3000/").unwrap();
assert_eq!(c.base_url, "http://localhost:3000");
}
#[test]
fn chat_request_serializes() {
let req = ChatRequest {
model: "test".into(),
messages: vec![Message {
role: "user".into(),
content: "hi".into(),
}],
max_tokens: Some(10),
..Default::default()
};
let json = serde_json::to_string(&req).unwrap();
assert!(json.contains("\"model\":\"test\""));
assert!(json.contains("\"max_tokens\":10"));
assert!(!json.contains("stream"));
}
#[test]
fn chat_response_deserializes() {
let json = r#"{
"id": "chatcmpl-test",
"choices": [{
"index": 0,
"message": {"role": "assistant", "content": "hello"},
"finish_reason": "stop"
}],
"usage": {
"prompt_tokens": 5,
"completion_tokens": 1,
"total_tokens": 6
}
}"#;
let resp: ChatResponse = serde_json::from_str(json).unwrap();
assert_eq!(resp.choices.len(), 1);
assert_eq!(resp.choices[0].message.content, "hello");
}
#[test]
fn models_response_deserializes() {
let json = r#"{
"object": "list",
"data": [
{"id": "llama-3-8b", "object": "model", "created": 0, "owned_by": "user"}
]
}"#;
let resp: ModelsResponse = serde_json::from_str(json).unwrap();
assert_eq!(resp.data.len(), 1);
assert_eq!(resp.data[0].id, "llama-3-8b");
}
#[test]
fn api_key_from_env() {
std::env::set_var("ARKNET_WALLET", "ark1fromenv");
let c = Client::new("http://localhost:1234").unwrap();
assert_eq!(c.api_key.as_deref(), Some("ark1fromenv"));
std::env::remove_var("ARKNET_WALLET");
}
#[test]
fn prefer_tee_serialized_when_set() {
let req = ChatRequest {
model: "test".into(),
messages: vec![],
prefer_tee: Some(true),
..Default::default()
};
let json = serde_json::to_string(&req).unwrap();
assert!(json.contains("\"prefer_tee\":true"));
}
#[test]
fn prefer_tee_omitted_when_none() {
let req = ChatRequest {
model: "test".into(),
messages: vec![],
..Default::default()
};
let json = serde_json::to_string(&req).unwrap();
assert!(!json.contains("prefer_tee"));
}
#[test]
fn require_https_serialized_when_set() {
let req = ChatRequest {
model: "test".into(),
messages: vec![],
require_https: Some(true),
..Default::default()
};
let json = serde_json::to_string(&req).unwrap();
assert!(json.contains("\"require_https\":true"));
}
#[test]
fn connect_options_defaults() {
let opts = ConnectOptions::default();
assert!(opts.seeds.is_empty());
assert!(!opts.require_https);
}
}