o402 0.1.5

OpenAI-compatible gateway, paid with x402.
//! Shared Axum state.

use std::sync::{Arc, Mutex};
use std::time::{Duration, Instant};

use hashbrown::HashMap;
use r402_server::ResourceServer;
use serde::Deserialize;
use url::Url;

use crate::config::{Config, ModelConfig, PaymentConfig, PricingConfig, UpstreamConfig};
use crate::error::AppError;
use crate::payment::inflight::InFlightSettles;
use crate::payment::price::{LoadedRates, Scheme};
use crate::payment::{schemes, tags};
use crate::proxy;

/// How long to reuse the first upstream's `GET /v1/models` for `/v1/pricing`.
const UPSTREAM_MODELS_TTL: Duration = Duration::from_secs(60);

/// Cached `GET /v1/models` rows: `(id, owned_by)`.
type UpstreamModelRows = Vec<(String, String)>;

/// Cloneable handle to process configuration and upstream clients.
#[derive(Clone, Debug)]
pub(crate) struct AppState {
    /// Shared inner state.
    inner: Arc<Inner>,
}

/// Process configuration, clients, and optional payment runtime.
#[derive(Debug)]
struct Inner {
    /// Validated configuration.
    config: Config,
    /// Clients keyed by `[[upstreams]]` name.
    clients: HashMap<String, reqwest::Client>,
    /// Resource server when payment is on.
    resource_server: Option<ResourceServer>,
    /// In-flight settle set (empty when payment is off).
    inflight: InFlightSettles,
    /// Integer rates keyed by model id (empty when payment is off).
    rates: HashMap<String, LoadedRates>,
    /// Default exact rates when payment is on.
    exact_default: Option<LoadedRates>,
    /// Default upto rates when payment is on.
    upto_default: Option<LoadedRates>,
    /// Cached first-upstream `/v1/models` ids for `/v1/pricing`.
    upstream_models: Mutex<Option<(Instant, UpstreamModelRows)>>,
}

impl AppState {
    /// Wrap `config`, build upstream clients, and (when paid) the resource server.
    ///
    /// # Errors
    ///
    /// Returns [`AppError::HttpClient`] when a client cannot be built, or
    /// [`AppError::Config`] when payment accepts cannot be turned into tags.
    pub(crate) fn new(config: Config) -> Result<Self, AppError> {
        let mut clients = HashMap::with_capacity(config.upstreams.len());
        for upstream in &config.upstreams {
            let mut builder = reqwest::Client::builder()
                .redirect(reqwest::redirect::Policy::none())
                .connect_timeout(Duration::from_secs(upstream.connect_timeout_secs))
                .timeout(Duration::from_secs(upstream.timeout_secs));
            if crate::config::is_loopback(&upstream.base_url) {
                builder = builder.no_proxy();
            }
            let client = builder.build().map_err(|source| AppError::HttpClient {
                name: upstream.name.clone(),
                source,
            })?;
            clients.insert(upstream.name.clone(), client);
        }

        let inflight = InFlightSettles::new();
        let (resource_server, rates, exact_default, upto_default) = if config.payment.enabled {
            let loaded =
                tags::validate_payment(&config.payment, &config.models, config.pricing.as_ref())?;
            let base_url = config.server.base_url.as_ref().ok_or_else(|| {
                crate::config::ConfigError::Validation(
                    "server.base_url is required when payment.enabled = true".to_owned(),
                )
            })?;
            (
                Some(schemes::resource_server(&config.payment, base_url)?),
                loaded.models,
                Some(loaded.exact_default),
                Some(loaded.upto_default),
            )
        } else {
            (None, HashMap::new(), None, None)
        };

        Ok(Self {
            inner: Arc::new(Inner {
                config,
                clients,
                resource_server,
                inflight,
                rates,
                exact_default,
                upto_default,
                upstream_models: Mutex::new(None),
            }),
        })
    }

    /// Whether x402 is enabled.
    #[must_use]
    pub(crate) fn payment_enabled(&self) -> bool {
        self.inner.config.payment.enabled
    }

    /// Payment config.
    #[must_use]
    pub(crate) fn payment(&self) -> &PaymentConfig {
        &self.inner.config.payment
    }

    /// Public origin used to build 402 `resource.url`.
    #[must_use]
    pub(crate) fn base_url(&self) -> Option<&Url> {
        self.inner.config.server.base_url.as_ref()
    }

    /// Resource server when payment is on.
    #[must_use]
    pub(crate) fn resource_server(&self) -> Option<&ResourceServer> {
        self.inner.resource_server.as_ref()
    }

    /// In-flight settle handle.
    #[must_use]
    pub(crate) fn inflight(&self) -> InFlightSettles {
        self.inner.inflight.clone()
    }

    /// Catalog in config order.
    #[must_use]
    pub(crate) fn models(&self) -> &[ModelConfig] {
        &self.inner.config.models
    }

    /// Look up a catalog model by client-facing id.
    #[must_use]
    pub(crate) fn model(&self, id: &str) -> Option<&ModelConfig> {
        self.inner.config.models.iter().find(|model| model.id == id)
    }

    /// Rates for `id` and `scheme`. Catalog override if it matches, else default.
    #[must_use]
    pub(crate) fn rates_for(&self, id: &str, scheme: Scheme) -> Option<&LoadedRates> {
        if let Some(loaded) = self.inner.rates.get(id)
            && loaded.scheme == scheme
        {
            return Some(loaded);
        }
        match scheme {
            Scheme::Exact => self.inner.exact_default.as_ref(),
            Scheme::Upto => self.inner.upto_default.as_ref(),
        }
    }

    /// Catalog model (if any) plus the upstream used to proxy `model_id`.
    #[must_use]
    pub(crate) fn resolve_upstream(
        &self,
        model_id: Option<&str>,
    ) -> Option<(&UpstreamConfig, &reqwest::Client, Option<&ModelConfig>)> {
        let catalog = model_id.and_then(|id| self.model(id));
        let name = catalog.map(|model| model.upstream.as_str()).or_else(|| {
            self.inner
                .config
                .upstreams
                .first()
                .map(|up| up.name.as_str())
        })?;
        let (upstream, client) = self.upstream(name)?;
        Some((upstream, client, catalog))
    }

    /// Pricing table, if present.
    #[must_use]
    pub(crate) fn pricing(&self) -> Option<&PricingConfig> {
        self.inner.config.pricing.as_ref()
    }

    /// Upstream config and HTTP client for `name`.
    #[must_use]
    pub(crate) fn upstream(&self, name: &str) -> Option<(&UpstreamConfig, &reqwest::Client)> {
        let upstream = self
            .inner
            .config
            .upstreams
            .iter()
            .find(|upstream| upstream.name == name)?;
        let client = self.inner.clients.get(name)?;
        Some((upstream, client))
    }

    /// First-upstream OpenAI model list `(id, owned_by)`, cached 60s.
    ///
    /// Empty when the fetch fails. Used by `GET /v1/pricing` so the demo
    /// catalog is not limited to `[[models]]`.
    pub(crate) async fn upstream_pricing_rows(&self) -> UpstreamModelRows {
        if let Some(rows) = self.cached_upstream_models() {
            return rows;
        }
        let rows = self.fetch_upstream_models().await.unwrap_or_default();
        if let Ok(mut guard) = self.inner.upstream_models.lock() {
            *guard = Some((Instant::now(), rows.clone()));
        }
        rows
    }

    #[allow(
        clippy::significant_drop_tightening,
        reason = "clone the cache entry then drop the mutex"
    )]
    fn cached_upstream_models(&self) -> Option<UpstreamModelRows> {
        let guard = self.inner.upstream_models.lock().ok()?;
        match guard.as_ref() {
            Some((at, rows)) if at.elapsed() < UPSTREAM_MODELS_TTL => Some(rows.clone()),
            _ => None,
        }
    }

    async fn fetch_upstream_models(&self) -> Option<UpstreamModelRows> {
        let (upstream, client, _) = self.resolve_upstream(None)?;
        let url = proxy::join_origin(&upstream.base_url, "/v1/models", None);
        let auth = format!("Bearer {}", upstream.api_key.expose());
        let response = client
            .get(url)
            .header(reqwest::header::AUTHORIZATION, auth)
            .send()
            .await
            .inspect_err(|error| tracing::warn!(error = %error, "upstream /v1/models failed"))
            .ok()?;
        if !response.status().is_success() {
            tracing::warn!(status = %response.status(), "upstream /v1/models not ok");
            return None;
        }
        let list: OpenAiModels = response
            .json()
            .await
            .inspect_err(|error| tracing::warn!(error = %error, "upstream /v1/models json"))
            .ok()?;
        Some(
            list.data
                .into_iter()
                .filter(|row| !row.id.is_empty())
                .map(|row| {
                    let owned = row
                        .owned_by
                        .filter(|value| !value.is_empty())
                        .unwrap_or_else(|| "system".to_owned());
                    (row.id, owned)
                })
                .collect(),
        )
    }
}

#[derive(Deserialize)]
struct OpenAiModels {
    #[serde(default)]
    data: Vec<OpenAiModelCard>,
}

#[derive(Deserialize)]
struct OpenAiModelCard {
    id: String,
    #[serde(default)]
    owned_by: Option<String>,
}