o402 0.1.2

OpenAI-compatible gateway, paid with x402.
//! [`TeeBody`] copies SSE to the client, parses usage, and spawns one settle.

use std::pin::Pin;
use std::sync::Arc;
use std::task::{Context, Poll};

use axum::body::Bytes;
use futures_util::Stream;
use r402_http::server::SettlementOverrides;
use r402_protocol::payment::{Extensions, PaymentRequirements};
use r402_server::{ResourceServer, SettlePhase, WirePaymentPayload};

use super::inflight::{InFlightSettles, SettleOnce};
use super::price::{self, LoadedRates};
use super::usage::{SseUsageParser, Usage};
use crate::config::UsagePolicy;

/// Owned inputs for the spawned `settle_payment` task.
#[derive(Debug)]
pub(super) struct SpawnSettle {
    pub(super) server: ResourceServer,
    pub(super) payload: WirePaymentPayload,
    pub(super) requirements: PaymentRequirements,
    pub(super) advertised: Extensions,
    pub(super) resource_url: String,
    pub(super) inflight: InFlightSettles,
    pub(super) signature: Vec<u8>,
    pub(super) overrides: Option<SettlementOverrides>,
}

/// Metering inputs plus the settle identity for one upto stream.
#[derive(Debug)]
pub(super) struct SettleJob {
    pub(super) server: ResourceServer,
    pub(super) payload: WirePaymentPayload,
    pub(super) requirements: PaymentRequirements,
    pub(super) advertised: Extensions,
    pub(super) resource_url: String,
    pub(super) inflight: InFlightSettles,
    pub(super) signature: Vec<u8>,
    pub(super) rates: LoadedRates,
    pub(super) ceiling: u128,
    pub(super) missing_usage: UsagePolicy,
    pub(super) abort_usage: UsagePolicy,
}

/// Upstream SSE body that tees bytes into [`SseUsageParser`].
pub(super) struct TeeBody {
    inner: Pin<Box<dyn Stream<Item = Result<Bytes, reqwest::Error>> + Send>>,
    parser: SseUsageParser,
    once: Arc<SettleOnce>,
    job: Option<SettleJob>,
}

impl std::fmt::Debug for TeeBody {
    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
        f.debug_struct("TeeBody")
            .field("usage", &self.parser.usage())
            .field("pending", &self.job.is_some())
            .finish_non_exhaustive()
    }
}

impl TeeBody {
    #[must_use]
    pub(super) fn new<S>(stream: S, job: SettleJob) -> Self
    where
        S: Stream<Item = Result<Bytes, reqwest::Error>> + Send + 'static,
    {
        Self {
            inner: Box::pin(stream),
            parser: SseUsageParser::new(),
            once: Arc::new(SettleOnce::new()),
            job: Some(job),
        }
    }

    fn take_job(&mut self) -> Option<SettleJob> {
        if self.once.take() {
            self.job.take()
        } else {
            None
        }
    }

    fn settle_complete(&mut self) {
        let Some(job) = self.take_job() else {
            return;
        };
        self.parser.finish();
        let policy = job.missing_usage;
        spawn_metered(job, self.parser.usage(), policy);
    }

    fn settle_abort(&mut self) {
        let Some(job) = self.take_job() else {
            return;
        };
        let policy = job.abort_usage;
        spawn_metered(job, None, policy);
    }
}

impl Stream for TeeBody {
    type Item = Result<Bytes, std::io::Error>;

    fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
        let this = self.get_mut();
        match this.inner.as_mut().poll_next(cx) {
            Poll::Pending => Poll::Pending,
            Poll::Ready(Some(Ok(chunk))) => {
                this.parser.push(&chunk);
                Poll::Ready(Some(Ok(chunk)))
            }
            Poll::Ready(Some(Err(error))) => {
                this.settle_complete();
                Poll::Ready(Some(Err(std::io::Error::other(error))))
            }
            Poll::Ready(None) => {
                this.settle_complete();
                Poll::Ready(None)
            }
        }
    }
}

impl Drop for TeeBody {
    fn drop(&mut self) {
        self.settle_abort();
    }
}

fn spawn_metered(job: SettleJob, usage: Option<Usage>, policy: UsagePolicy) {
    let actual = price::actual(&job.rates, job.ceiling, usage.as_ref(), policy);
    spawn_settle_task(SpawnSettle {
        server: job.server,
        payload: job.payload,
        requirements: job.requirements,
        advertised: job.advertised,
        resource_url: job.resource_url,
        inflight: job.inflight,
        signature: job.signature,
        overrides: Some(SettlementOverrides::amount(actual.to_string())),
    });
}

/// Increments the spawned counter, runs `settle_payment`, then drops the hash.
pub(super) fn spawn_settle_task(job: SpawnSettle) {
    job.inflight.increment_spawned();
    drop(tokio::spawn(async move {
        let result = job
            .server
            .settle_payment(
                &job.payload,
                &job.requirements,
                job.overrides.as_ref(),
                SettlePhase::AfterHandler,
                Some(job.resource_url.as_str()),
                Some(&job.advertised),
            )
            .await;
        match result {
            Ok(ref settlement) if settlement.is_success() => {
                tracing::info!(success = true, "settle completed");
            }
            Ok(_) => tracing::warn!(success = false, "settle completed"),
            Err(ref error) => tracing::warn!(error = %error, "settle failed"),
        }
        job.inflight.remove(&job.signature);
        job.inflight.decrement_spawned();
    }));
}