use serde_json::{json, Value};
#[derive(Debug, Clone, PartialEq)]
pub struct InferArgs {
pub model: crate::model_uri::ModelAddress,
pub prompt: String,
pub system: Option<String>,
pub max_tokens: Option<u64>,
pub broker: Option<String>,
}
#[derive(Debug, Clone, PartialEq)]
pub struct ChatArgs {
pub model: crate::model_uri::ModelAddress,
pub system: Option<String>,
pub broker: Option<String>,
}
pub fn parse_infer_args(args: &[String]) -> Result<InferArgs, String> {
let mut model: Option<crate::model_uri::ModelAddress> = None;
let mut prompt: Option<String> = None;
let mut system: Option<String> = None;
let mut max_tokens: Option<u64> = None;
let mut broker: Option<String> = None;
let mut it = args.iter().peekable();
while let Some(arg) = it.next() {
match arg.as_str() {
"-m" | "--prompt" => {
let v = it.next().ok_or_else(|| format!("{arg} requires a value"))?;
prompt = Some(v.clone());
}
"--max-tokens" => {
let v = it
.next()
.ok_or_else(|| "--max-tokens requires a value".to_string())?;
max_tokens = Some(
v.parse::<u64>()
.map_err(|_| format!("invalid --max-tokens value '{v}'"))?,
);
}
"--system" => {
let v = it
.next()
.ok_or_else(|| "--system requires a value".to_string())?;
system = Some(v.clone());
}
"--broker" => {
let v = it
.next()
.ok_or_else(|| "--broker requires a value".to_string())?;
broker = Some(v.clone());
}
other => {
if model.is_some() {
return Err(format!("unexpected argument '{other}'"));
}
let addr = crate::model_uri::parse_model_address(other)
.ok_or_else(|| format!("not a model address: '{other}'"))?;
model = Some(addr);
}
}
}
let model = model.ok_or_else(|| {
"zc infer requires a zc://<owner>/<name> or zc://<uuid> argument".to_string()
})?;
let prompt = prompt.ok_or_else(|| "zc infer requires -m/--prompt <text>".to_string())?;
Ok(InferArgs {
model,
prompt,
system,
max_tokens,
broker,
})
}
pub fn parse_chat_args(args: &[String]) -> Result<ChatArgs, String> {
let mut model: Option<crate::model_uri::ModelAddress> = None;
let mut system: Option<String> = None;
let mut broker: Option<String> = None;
let mut it = args.iter().peekable();
while let Some(arg) = it.next() {
match arg.as_str() {
"--system" => {
let v = it
.next()
.ok_or_else(|| "--system requires a value".to_string())?;
system = Some(v.clone());
}
"--broker" => {
let v = it
.next()
.ok_or_else(|| "--broker requires a value".to_string())?;
broker = Some(v.clone());
}
other => {
if model.is_some() {
return Err(format!("unexpected argument '{other}'"));
}
let addr = crate::model_uri::parse_model_address(other)
.ok_or_else(|| format!("not a model address: '{other}'"))?;
model = Some(addr);
}
}
}
let model = model.ok_or_else(|| {
"zc chat requires a zc://<owner>/<name> or zc://<uuid> argument".to_string()
})?;
Ok(ChatArgs {
model,
system,
broker,
})
}
pub fn build_messages(prompt: &str, system: Option<&str>) -> Vec<Value> {
let mut messages = Vec::new();
if let Some(sys) = system {
messages.push(json!({"role": "system", "content": sys}));
}
messages.push(json!({"role": "user", "content": prompt}));
messages
}
pub fn build_infer_body(
model_uuid: &str,
prompt: &str,
system: Option<&str>,
max_tokens: Option<u64>,
) -> Value {
build_chat_body(model_uuid, &build_messages(prompt, system), max_tokens)
}
pub fn build_chat_body(model_uuid: &str, messages: &[Value], max_tokens: Option<u64>) -> Value {
let mut body = json!({
"model": format!("zc://{model_uuid}"),
"messages": messages,
});
if let Some(mt) = max_tokens {
body["max_tokens"] = json!(mt);
}
body
}
pub fn format_infer_output(resp: &Value) -> String {
let content = resp.get("content").and_then(|v| v.as_str()).unwrap_or("");
let total_tokens = resp
.get("usage")
.and_then(|u| u.get("total_tokens"))
.and_then(|v| v.as_u64())
.unwrap_or(0);
let charged_zkcr = resp
.get("charged_zkcr")
.and_then(|v| v.as_f64())
.unwrap_or(0.0);
let provider = resp
.get("provider")
.and_then(|v| v.as_str())
.unwrap_or("unknown");
format!("{content}\n\n{total_tokens} tokens -> {charged_zkcr} zkcr (provider: {provider})")
}
pub fn format_infer_error(body: &str) -> String {
match serde_json::from_str::<Value>(body) {
Ok(v) => v
.get("error")
.and_then(|e| e.as_str())
.map(|s| s.to_string())
.unwrap_or_else(|| body.to_string()),
Err(_) => body.to_string(),
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::model_uri::ModelAddress;
const UUID: &str = "ae5f3db4-437a-40d1-93ec-c8258315d69a";
#[test]
fn parses_minimal_infer_args() {
let args = vec![
format!("zc://{UUID}"),
"-m".to_string(),
"hello".to_string(),
];
let parsed = parse_infer_args(&args).unwrap();
assert_eq!(parsed.model, ModelAddress::Uuid(UUID.to_string()));
assert_eq!(parsed.prompt, "hello");
assert_eq!(parsed.system, None);
assert_eq!(parsed.max_tokens, None);
assert_eq!(parsed.broker, None);
}
#[test]
fn parses_full_infer_args() {
let args = vec![
format!("zc://model-{UUID}"),
"--prompt".to_string(),
"hi there".to_string(),
"--max-tokens".to_string(),
"256".to_string(),
"--system".to_string(),
"be terse".to_string(),
"--broker".to_string(),
"http://localhost:9000".to_string(),
];
let parsed = parse_infer_args(&args).unwrap();
assert_eq!(parsed.model, ModelAddress::Uuid(UUID.to_string()));
assert_eq!(parsed.prompt, "hi there");
assert_eq!(parsed.system.as_deref(), Some("be terse"));
assert_eq!(parsed.max_tokens, Some(256));
assert_eq!(parsed.broker.as_deref(), Some("http://localhost:9000"));
}
#[test]
fn rejects_invalid_model_uri() {
let args = vec![
"zc://node-abc123".to_string(),
"-m".to_string(),
"hi".to_string(),
];
let err = parse_infer_args(&args).unwrap_err();
assert!(err.contains("not a model address"));
}
#[test]
fn rejects_missing_prompt() {
let args = vec![format!("zc://{UUID}")];
let err = parse_infer_args(&args).unwrap_err();
assert!(err.contains("--prompt"));
}
#[test]
fn rejects_missing_model() {
let args = vec!["-m".to_string(), "hi".to_string()];
let err = parse_infer_args(&args).unwrap_err();
assert!(err.contains("<owner>/<name>"), "{err}");
assert!(err.contains("zc://"), "{err}");
}
#[test]
fn parses_chat_args() {
let args = vec![
format!("zc://{UUID}"),
"--system".to_string(),
"be terse".to_string(),
];
let parsed = parse_chat_args(&args).unwrap();
assert_eq!(parsed.model, ModelAddress::Uuid(UUID.to_string()));
assert_eq!(parsed.system.as_deref(), Some("be terse"));
assert_eq!(parsed.broker, None);
}
#[test]
fn build_infer_body_without_system_or_max_tokens() {
let body = build_infer_body(UUID, "hello", None, None);
assert_eq!(body["model"], json!(format!("zc://{UUID}")));
assert_eq!(
body["messages"],
json!([{"role": "user", "content": "hello"}])
);
assert!(body.get("max_tokens").is_none());
}
#[test]
fn build_infer_body_with_system_and_max_tokens() {
let body = build_infer_body(UUID, "hello", Some("be terse"), Some(64));
assert_eq!(
body["messages"],
json!([
{"role": "system", "content": "be terse"},
{"role": "user", "content": "hello"}
])
);
assert_eq!(body["max_tokens"], json!(64));
}
#[test]
fn format_infer_output_renders_content_and_usage() {
let resp = json!({
"content": "Paris is the capital of France.",
"usage": {"prompt_tokens": 10, "completion_tokens": 8, "total_tokens": 18},
"provider": "specialized",
"charged_zkcr": 0.0042
});
let out = format_infer_output(&resp);
assert!(out.starts_with("Paris is the capital of France."));
assert!(out.contains("18 tokens -> 0.0042 zkcr (provider: specialized)"));
}
#[test]
fn format_infer_output_defaults_on_missing_fields() {
let resp = json!({});
let out = format_infer_output(&resp);
assert!(out.contains("0 tokens -> 0 zkcr (provider: unknown)"));
}
#[test]
fn format_infer_error_extracts_error_field() {
let body = r#"{"error": "model not found"}"#;
assert_eq!(format_infer_error(body), "model not found");
}
#[test]
fn format_infer_error_falls_back_to_raw_body() {
let body = "not json";
assert_eq!(format_infer_error(body), "not json");
}
}