use std::future::Future;
use std::sync::Arc;
use std::sync::atomic::{AtomicU64, Ordering};
use parking_lot::Mutex;
use nexil::{CancellationToken, Tool};
use serde_json::Value;
pub type WrapToolsFn = Arc<dyn Fn(Vec<Tool>) -> Vec<Tool> + Send + Sync>;
pub type DispatchFn = Arc<dyn Fn(Value) -> futures::future::BoxFuture<'static, ()> + Send + Sync>;
tokio::task_local! {
static TURN_CTX: TurnContext;
}
#[derive(Clone, Default)]
pub struct TurnUsage {
input: Arc<AtomicU64>,
output: Arc<AtomicU64>,
}
impl TurnUsage {
pub fn record(&self, input_tokens: u64, output_tokens: u64) {
self.input.fetch_add(input_tokens, Ordering::Relaxed);
self.output.fetch_add(output_tokens, Ordering::Relaxed);
}
pub fn input_tokens(&self) -> u64 {
self.input.load(Ordering::Relaxed)
}
pub fn output_tokens(&self) -> u64 {
self.output.load(Ordering::Relaxed)
}
pub fn total_tokens(&self) -> u64 {
self.input_tokens() + self.output_tokens()
}
}
#[derive(Debug, Clone)]
pub struct OutboundMedia {
pub path: String,
pub media_type: String,
pub mime_type: String,
}
pub struct TurnContext {
pub cancellation: CancellationToken,
pub wrap_tools: Option<WrapToolsFn>,
pub usage: TurnUsage,
pub save_events: Arc<Mutex<Vec<(String, Value)>>>,
pub dispatch: Option<DispatchFn>,
pub outbound_media: Arc<Mutex<Vec<OutboundMedia>>>,
}
pub async fn with_turn_context<F, T>(ctx: TurnContext, fut: F) -> T
where
F: Future<Output = T>,
{
TURN_CTX.scope(ctx, fut).await
}
pub fn turn_cancellation() -> Option<CancellationToken> {
TURN_CTX.try_with(|ctx| ctx.cancellation.clone()).ok()
}
pub fn turn_usage() -> Option<TurnUsage> {
TURN_CTX.try_with(|ctx| ctx.usage.clone()).ok()
}
pub fn record_turn_usage(input_tokens: u64, output_tokens: u64) {
let _ = TURN_CTX.try_with(|ctx| ctx.usage.record(input_tokens, output_tokens));
}
pub fn push_save_event(name: &str, data: Value) {
let _ = TURN_CTX.try_with(|ctx| {
ctx.save_events.lock().push((name.to_owned(), data));
});
}
pub fn drain_save_events() -> Vec<(String, Value)> {
TURN_CTX
.try_with(|ctx| std::mem::take(&mut *ctx.save_events.lock()))
.unwrap_or_default()
}
pub async fn dispatch_mid_turn(envelope: Value) {
let dispatch = TURN_CTX.try_with(|ctx| ctx.dispatch.clone()).ok().flatten();
if let Some(f) = dispatch {
f(envelope).await;
}
}
pub fn push_outbound_media(media: OutboundMedia) {
let _ = TURN_CTX.try_with(|ctx| {
ctx.outbound_media.lock().push(media);
});
}
pub fn drain_outbound_media() -> Vec<OutboundMedia> {
TURN_CTX
.try_with(|ctx| std::mem::take(&mut *ctx.outbound_media.lock()))
.unwrap_or_default()
}
pub fn mime_from_extension(path: &std::path::Path) -> &'static str {
match path.extension().and_then(|e| e.to_str()) {
Some("png") => "image/png",
Some("jpg" | "jpeg") => "image/jpeg",
Some("gif") => "image/gif",
Some("webp") => "image/webp",
Some("svg") => "image/svg+xml",
Some("mp4") => "video/mp4",
Some("mp3") => "audio/mpeg",
Some("ogg") => "audio/ogg",
Some("pdf") => "application/pdf",
_ => "application/octet-stream",
}
}
pub fn media_type_from_mime(mime: &str) -> &'static str {
if mime.starts_with("image/") {
"image"
} else if mime.starts_with("video/") {
"video"
} else if mime.starts_with("audio/") {
"audio"
} else {
"document"
}
}
pub fn turn_wrap_tools() -> Option<WrapToolsFn> {
TURN_CTX
.try_with(|ctx| ctx.wrap_tools.clone())
.ok()
.flatten()
}
pub type InjectInboundFn =
Arc<dyn Fn(Value) -> futures::future::BoxFuture<'static, ()> + Send + Sync>;
static INBOUND_INJECTOR: std::sync::LazyLock<Mutex<Option<InjectInboundFn>>> =
std::sync::LazyLock::new(|| Mutex::new(None));
pub fn set_inbound_injector(f: InjectInboundFn) {
*INBOUND_INJECTOR.lock() = Some(f);
}
pub fn inbound_injector() -> Option<InjectInboundFn> {
INBOUND_INJECTOR.lock().clone()
}
pub async fn inject_inbound(envelope: Value) {
if let Some(f) = inbound_injector() {
f(envelope).await;
}
}
pub struct BudgetLedger {
remaining: AtomicU64,
consumed: AtomicU64,
}
impl BudgetLedger {
pub fn new() -> Self {
Self {
remaining: AtomicU64::new(u64::MAX),
consumed: AtomicU64::new(0),
}
}
pub fn with_budget(max_tokens: u64) -> Self {
Self {
remaining: AtomicU64::new(max_tokens),
consumed: AtomicU64::new(0),
}
}
pub fn try_spend(&self, amount: u64) -> bool {
loop {
let current = self.remaining.load(Ordering::Acquire);
if current < amount {
return false;
}
if self
.remaining
.compare_exchange_weak(
current,
current - amount,
Ordering::AcqRel,
Ordering::Acquire,
)
.is_ok()
{
self.consumed.fetch_add(amount, Ordering::Relaxed);
return true;
}
}
}
pub fn remaining(&self) -> u64 {
self.remaining.load(Ordering::Relaxed)
}
pub fn consumed(&self) -> u64 {
self.consumed.load(Ordering::Relaxed)
}
}
impl Default for BudgetLedger {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn budget_unlimited_by_default() {
let b = BudgetLedger::new();
assert!(b.try_spend(1_000_000));
assert_eq!(b.consumed(), 1_000_000);
}
#[test]
fn budget_rejects_overspend() {
let b = BudgetLedger::with_budget(100);
assert!(b.try_spend(60));
assert!(!b.try_spend(60));
assert_eq!(b.remaining(), 40);
assert_eq!(b.consumed(), 60);
}
#[test]
fn budget_concurrent_spend() {
use std::sync::Arc;
let b = Arc::new(BudgetLedger::with_budget(1000));
let handles: Vec<_> = (0..100)
.map(|_| {
let b = Arc::clone(&b);
std::thread::spawn(move || b.try_spend(10))
})
.collect();
let successes: usize = handles
.into_iter()
.map(|h| h.join().unwrap())
.filter(|&ok| ok)
.count();
assert_eq!(successes, 100);
assert_eq!(b.remaining(), 0);
assert_eq!(b.consumed(), 1000);
}
#[tokio::test]
async fn turn_context_propagates() {
let token = CancellationToken::new();
let token2 = token.clone();
let ctx = TurnContext {
cancellation: token,
wrap_tools: None,
usage: Default::default(),
save_events: Default::default(),
dispatch: None,
outbound_media: Default::default(),
};
let result = with_turn_context(ctx, async {
let t = turn_cancellation().unwrap();
assert!(!t.is_cancelled());
token2.cancel();
assert!(t.is_cancelled());
true
})
.await;
assert!(result);
}
#[tokio::test]
async fn turn_context_absent_returns_none() {
assert!(turn_cancellation().is_none());
assert!(turn_wrap_tools().is_none());
}
#[tokio::test]
async fn outbound_media_push_and_drain() {
let ctx = TurnContext {
cancellation: CancellationToken::new(),
wrap_tools: None,
usage: Default::default(),
save_events: Default::default(),
dispatch: None,
outbound_media: Default::default(),
};
with_turn_context(ctx, async {
push_outbound_media(OutboundMedia {
path: "/tmp/a.png".into(),
media_type: "image".into(),
mime_type: "image/png".into(),
});
push_outbound_media(OutboundMedia {
path: "/tmp/b.mp4".into(),
media_type: "video".into(),
mime_type: "video/mp4".into(),
});
let drained = drain_outbound_media();
assert_eq!(drained.len(), 2);
assert_eq!(drained[0].path, "/tmp/a.png");
assert_eq!(drained[1].path, "/tmp/b.mp4");
assert!(drain_outbound_media().is_empty());
})
.await;
}
#[tokio::test]
async fn outbound_media_drain_without_context_returns_empty() {
assert!(drain_outbound_media().is_empty());
}
#[tokio::test]
async fn inject_inbound_noop_without_injector() {
inject_inbound(serde_json::json!({"content": "test"})).await;
}
#[test]
fn mime_from_extension_known_types() {
assert_eq!(
mime_from_extension(std::path::Path::new("/tmp/photo.png")),
"image/png"
);
assert_eq!(
mime_from_extension(std::path::Path::new("file.jpg")),
"image/jpeg"
);
assert_eq!(
mime_from_extension(std::path::Path::new("video.mp4")),
"video/mp4"
);
assert_eq!(
mime_from_extension(std::path::Path::new("doc.pdf")),
"application/pdf"
);
}
#[test]
fn mime_from_extension_unknown_fallback() {
assert_eq!(
mime_from_extension(std::path::Path::new("file.xyz")),
"application/octet-stream"
);
assert_eq!(
mime_from_extension(std::path::Path::new("noext")),
"application/octet-stream"
);
}
#[test]
fn media_type_from_mime_categorizes() {
assert_eq!(media_type_from_mime("image/png"), "image");
assert_eq!(media_type_from_mime("image/jpeg"), "image");
assert_eq!(media_type_from_mime("video/mp4"), "video");
assert_eq!(media_type_from_mime("audio/mpeg"), "audio");
assert_eq!(media_type_from_mime("application/pdf"), "document");
assert_eq!(media_type_from_mime("application/octet-stream"), "document");
}
}