use std::sync::Arc;
use std::time::Duration;
use hashbrown::HashMap;
use r402_server::ResourceServer;
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};
#[derive(Clone, Debug)]
pub(crate) struct AppState {
inner: Arc<Inner>,
}
#[derive(Debug)]
struct Inner {
config: Config,
clients: HashMap<String, reqwest::Client>,
resource_server: Option<ResourceServer>,
inflight: InFlightSettles,
rates: HashMap<String, LoadedRates>,
exact_default: Option<LoadedRates>,
upto_default: Option<LoadedRates>,
}
impl AppState {
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,
}),
})
}
#[must_use]
pub(crate) fn payment_enabled(&self) -> bool {
self.inner.config.payment.enabled
}
#[must_use]
pub(crate) fn payment(&self) -> &PaymentConfig {
&self.inner.config.payment
}
#[must_use]
pub(crate) fn base_url(&self) -> Option<&Url> {
self.inner.config.server.base_url.as_ref()
}
#[must_use]
pub(crate) fn resource_server(&self) -> Option<&ResourceServer> {
self.inner.resource_server.as_ref()
}
#[must_use]
pub(crate) fn inflight(&self) -> InFlightSettles {
self.inner.inflight.clone()
}
#[must_use]
pub(crate) fn models(&self) -> &[ModelConfig] {
&self.inner.config.models
}
#[must_use]
pub(crate) fn model(&self, id: &str) -> Option<&ModelConfig> {
self.inner.config.models.iter().find(|model| model.id == id)
}
#[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(),
}
}
#[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))
}
#[must_use]
pub(crate) fn pricing(&self) -> Option<&PricingConfig> {
self.inner.config.pricing.as_ref()
}
#[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))
}
}