malvin 0.2.9

Non-interactive research and coding agent
use std::collections::HashMap;

use pi::provider::ModelCost;
use pi::sdk::{Config, ModelRegistry};
use serde::Deserialize;
use serde_json::{Map, Value};

#[path = "portkey_pricing_fetch.rs"]
mod fetch;

#[derive(Debug, Deserialize)]
struct PricingBody {
    pay_as_you_go: Option<PayGo>,
}

#[derive(Debug, Deserialize)]
struct PayGo {
    #[serde(rename = "request_token")]
    request: Option<PricePoint>,
    #[serde(rename = "response_token")]
    response: Option<PricePoint>,
    #[serde(rename = "cache_write_input_token")]
    cache_write: Option<PricePoint>,
    #[serde(rename = "cache_read_input_token")]
    cache_read: Option<PricePoint>,
}

#[derive(Debug, Deserialize)]
struct PricePoint {
    price: f64,
}

pub(super) fn host_is_portkey(base_url: &str) -> bool {
    let rest = base_url
        .trim()
        .trim_start_matches("https://")
        .trim_start_matches("http://");
    let host = rest.split(['/', '?', '#']).next().unwrap_or("");
    let host = host.split('@').next_back().unwrap_or(host);
    let host = host.split(':').next().unwrap_or(host);
    host.eq_ignore_ascii_case("api.portkey.ai")
}

pub(super) fn name_is_portkey_header(name: &str) -> bool {
    name.to_ascii_lowercase().starts_with("x-portkey-")
}

fn bare_model_id(model_id: &str) -> &str {
    let id = model_id.trim();
    let Some(rest) = id.strip_prefix('@') else {
        return id;
    };
    rest.split_once('/').map_or(id, |(_, model)| model)
}

fn header_provider(headers: &HashMap<String, String>) -> Option<String> {
    let value = headers
        .iter()
        .find(|(name, _)| name.eq_ignore_ascii_case("x-portkey-provider"))?
        .1
        .trim();
    if value.is_empty() || value.starts_with(['@', '$', '!']) {
        return None;
    }
    Some(value.to_string())
}

fn point_rate(point: Option<&PricePoint>) -> f64 {
    point.map_or(0.0, |p| p.price * 10_000.0)
}

pub(super) fn model_cost_from_body(body: &str) -> Option<ModelCost> {
    let parsed: PricingBody = serde_json::from_str(body).ok()?;
    let pay = parsed.pay_as_you_go?;
    let cost = ModelCost {
        input: point_rate(pay.request.as_ref()),
        output: point_rate(pay.response.as_ref()),
        cache_read: point_rate(pay.cache_read.as_ref()),
        cache_write: point_rate(pay.cache_write.as_ref()),
    };
    cost.is_priced().then_some(cost)
}

fn load_registry() -> Option<ModelRegistry> {
    let auth = pi::auth::AuthStorage::load(Config::auth_path()).ok()?;
    let path = pi::models::default_models_path(&Config::global_dir());
    Some(ModelRegistry::load(&auth, Some(path)))
}

fn entry_header_provider(entry: &pi::sdk::ModelEntry) -> Option<String> {
    header_provider(&entry.headers).or_else(|| header_provider(&entry.model.headers))
}

fn entry_is_portkey(entry: &pi::sdk::ModelEntry) -> bool {
    if host_is_portkey(&entry.model.base_url) {
        return true;
    }
    entry
        .headers
        .keys()
        .any(|name| name_is_portkey_header(name))
        || entry
            .model
            .headers
            .keys()
            .any(|name| name_is_portkey_header(name))
}

fn upstream_name(pi_provider: &str, from_header: Option<String>) -> Option<String> {
    if let Some(name) = from_header {
        return Some(name);
    }
    if pi_provider.eq_ignore_ascii_case("portkey") {
        None
    } else {
        Some(pi_provider.to_string())
    }
}

pub(super) fn rates_for_pi_model(provider: &str, model_id: &str) -> Option<ModelCost> {
    let registry = load_registry()?;
    let entry = registry.find(provider, model_id)?;
    if !entry_is_portkey(&entry) {
        return None;
    }
    let upstream = upstream_name(provider, entry_header_provider(&entry))?;
    fetch::lookup_rates(&upstream, bare_model_id(model_id))
}

fn token_count(obj: &Map<String, Value>, keys: &[&str]) -> u64 {
    keys.iter()
        .find_map(|key| obj.get(*key).and_then(Value::as_u64))
        .unwrap_or(0)
}

pub(crate) fn apply_portkey_cost_usd(provider: &str, model: &str, usage: &mut Value) {
    let Some(rates) = rates_for_pi_model(provider, model) else {
        return;
    };
    let Some(obj) = usage.as_object() else {
        return;
    };
    let usage_tokens = pi::model::Usage {
        input: token_count(obj, &["inputTokens", "input"]),
        output: token_count(obj, &["outputTokens", "output"]),
        cache_read: token_count(obj, &["cacheReadTokens", "cacheRead", "cache_read"]),
        cache_write: token_count(obj, &["cacheWriteTokens", "cacheWrite", "cache_write"]),
        ..pi::model::Usage::default()
    };
    let cost = super::usage_cost::cost_from_model_rates(&rates, &usage_tokens);
    if let Some(map) = usage.as_object_mut() {
        map.insert(
            "costUsd".into(),
            serde_json::json!({
                "input": cost.input,
                "output": cost.output,
                "cacheRead": cost.cache_read,
                "cacheWrite": cost.cache_write,
                "total": cost.total,
            }),
        );
    }
}

#[cfg(test)]
#[path = "portkey_pricing_tests.rs"]
mod portkey_pricing_tests;