use std::collections::HashMap;
use std::sync::Arc;
use serde_json::json;
use tinyagents::TinyAgentsError;
use tinyagents::harness::context::{RunConfig, RunContext};
use tinyagents::harness::events::AgentEvent;
use tinyagents::harness::message::{AssistantMessage, ContentBlock, Message};
use tinyagents::harness::middleware::{
BudgetLimits, BudgetMiddleware, BudgetTracker, MiddlewareStack,
};
use tinyagents::harness::model::{
ModelRequest, ModelResolutionSource, ModelResponse, ResolvedModel,
};
use tinyagents::harness::providers::MockModel;
use tinyagents::harness::runtime::AgentHarness;
use tinyagents::harness::testkit::{EventRecorder, FakeTool};
use tinyagents::harness::tool::ToolCall;
use tinyagents::harness::usage::Usage;
use tinyagents::registry::catalog::ModelPricing;
fn tool_call_response(id: &str, name: &str, input: u64, output: u64) -> ModelResponse {
ModelResponse {
message: AssistantMessage {
id: Some(format!("msg-{id}")),
content: Vec::new(),
tool_calls: vec![ToolCall::new(id, name, json!({}))],
usage: Some(Usage::new(input, output)),
},
usage: Some(Usage::new(input, output)),
finish_reason: Some("tool_calls".into()),
raw: None,
resolved_model: None,
}
}
fn text_response(text: &str, input: u64, output: u64) -> ModelResponse {
ModelResponse {
message: AssistantMessage {
id: None,
content: vec![ContentBlock::Text(text.into())],
tool_calls: Vec::new(),
usage: Some(Usage::new(input, output)),
},
usage: Some(Usage::new(input, output)),
finish_reason: Some("stop".into()),
raw: None,
resolved_model: None,
}
}
fn any_event(recorder: &EventRecorder, pred: impl Fn(&AgentEvent) -> bool) -> bool {
recorder.events().iter().any(pred)
}
#[tokio::test]
async fn token_budget_blocks_multi_call_run() {
let recorder = EventRecorder::new();
let model = MockModel::with_responses(vec![
tool_call_response("c1", "noop", 8, 0),
tool_call_response("c2", "noop", 8, 0),
tool_call_response("c3", "noop", 8, 0),
]);
let limits = BudgetLimits {
max_total_tokens: Some(20),
warn_fraction: Some(0.5),
..Default::default()
};
let mw = BudgetMiddleware::new(limits);
let tracker = mw.tracker();
let mut harness: AgentHarness<()> = AgentHarness::new();
harness
.register_model("mock", Arc::new(model))
.set_default_model("mock")
.register_tool(Arc::new(FakeTool::returning("noop", "ok")))
.push_middleware(Arc::new(mw));
let ctx = RunContext::new(RunConfig::new("budget-tokens"), ()).with_events(recorder.sink());
let err = harness
.invoke_in_context(&(), ctx, vec![Message::user("go")])
.await
.expect_err("the accumulated token budget must block the run");
assert!(
matches!(err, TinyAgentsError::LimitExceeded(_)),
"expected LimitExceeded, got {err:?}"
);
assert!(
any_event(&recorder, |e| matches!(e, AgentEvent::BudgetWarning { .. })),
"a BudgetWarning must be emitted as usage crosses warn_fraction"
);
assert!(
any_event(&recorder, |e| matches!(
e,
AgentEvent::BudgetExceeded { blocked: true, .. }
)),
"the blocking preflight must emit BudgetExceeded {{ blocked: true }}"
);
assert_eq!(
tracker.snapshot().usage.usage.effective_total(),
24,
"three recorded calls of 8 tokens each"
);
assert_eq!(
tracker.snapshot().usage.calls,
3,
"three calls were recorded"
);
}
#[tokio::test]
async fn shared_tracker_rolls_up_and_blocks_across_runs() {
let tracker = BudgetTracker::new();
let limits = BudgetLimits {
max_total_tokens: Some(30),
..Default::default()
};
let model_a = MockModel::with_responses(vec![
tool_call_response("a1", "noop", 8, 0),
text_response("done", 8, 0),
]);
let mut harness_a: AgentHarness<()> = AgentHarness::new();
harness_a
.register_model("mock", Arc::new(model_a))
.set_default_model("mock")
.register_tool(Arc::new(FakeTool::returning("noop", "ok")))
.push_middleware(Arc::new(
BudgetMiddleware::new(limits).with_tracker(tracker.clone()),
));
let run_a = harness_a
.invoke_in_context(
&(),
RunContext::new(RunConfig::new("budget-parent"), ()),
vec![Message::user("parent")],
)
.await
.expect("the first run stays under the shared budget");
assert!(run_a.final_response.is_some(), "first run completed");
assert_eq!(
tracker.snapshot().usage.usage.effective_total(),
16,
"the shared tracker holds the first run's spend"
);
let model_b = MockModel::with_responses(vec![
tool_call_response("b1", "noop", 8, 0),
tool_call_response("b2", "noop", 8, 0),
]);
let mut harness_b: AgentHarness<()> = AgentHarness::new();
harness_b
.register_model("mock", Arc::new(model_b))
.set_default_model("mock")
.register_tool(Arc::new(FakeTool::returning("noop", "ok")))
.push_middleware(Arc::new(
BudgetMiddleware::new(limits).with_tracker(tracker.clone()),
));
let err = harness_b
.invoke_in_context(
&(),
RunContext::new(RunConfig::new("budget-child"), ()),
vec![Message::user("child")],
)
.await
.expect_err("the shared budget must block the second run");
assert!(
matches!(err, TinyAgentsError::LimitExceeded(_)),
"expected LimitExceeded, got {err:?}"
);
assert_eq!(
tracker.snapshot().usage.usage.effective_total(),
32,
"spend from both runs accumulates in the shared tracker"
);
assert_eq!(
tracker.snapshot().usage.calls,
4,
"four calls total across runs"
);
}
#[tokio::test]
async fn cost_pricing_records_and_enforces_money_budget() {
let recorder = EventRecorder::new();
let mut ctx = RunContext::new(RunConfig::new("budget-cost"), ()).with_events(recorder.sink());
let mut pricing: HashMap<String, ModelPricing> = HashMap::new();
pricing.insert(
"m".to_string(),
ModelPricing {
input_per_token: Some(1.0),
output_per_token: Some(1.0),
..Default::default()
},
);
let mw = BudgetMiddleware::new(BudgetLimits {
max_cost: Some(5.0),
..Default::default()
})
.with_pricing(pricing);
let tracker = mw.tracker();
let mut stack: MiddlewareStack<()> = MiddlewareStack::new();
stack.push(Arc::new(mw));
let mut resp = ModelResponse {
message: AssistantMessage {
id: None,
content: vec![ContentBlock::Text("priced".into())],
tool_calls: Vec::new(),
usage: Some(Usage::new(4, 2)),
},
usage: Some(Usage::new(4, 2)),
finish_reason: Some("stop".into()),
raw: None,
resolved_model: Some(ResolvedModel {
name: "m".into(),
requested: None,
source: ModelResolutionSource::RegistryDefault,
}),
};
stack
.run_after_model(&mut ctx, &(), &mut resp)
.await
.expect("recording spend does not fail");
assert!(
(tracker.snapshot().cost.total_cost - 6.0).abs() < 1e-9,
"4+2 tokens at 1.0/token should cost 6.0, got {}",
tracker.snapshot().cost.total_cost
);
assert!(
any_event(&recorder, |e| matches!(e, AgentEvent::CostRecorded { .. })),
"a CostRecorded event must be emitted when priced cost is positive"
);
assert!(
any_event(&recorder, |e| matches!(e, AgentEvent::UsageRecorded { .. })),
"a UsageRecorded event accompanies the recorded usage"
);
let mut req = ModelRequest::new(vec![Message::user("go")]);
let err = stack
.run_before_model(&mut ctx, &(), &mut req)
.await
.expect_err("the cost budget must block the next model call");
assert!(
matches!(err, TinyAgentsError::LimitExceeded(_)),
"expected LimitExceeded, got {err:?}"
);
assert!(
any_event(&recorder, |e| matches!(
e,
AgentEvent::BudgetExceeded { blocked: true, .. }
)),
"the blocking preflight emits BudgetExceeded {{ blocked: true }}"
);
}
#[tokio::test]
async fn below_threshold_run_completes_without_exceeding() {
let recorder = EventRecorder::new();
let model = MockModel::with_responses(vec![
tool_call_response("s1", "noop", 5, 0),
text_response("all good", 5, 0),
]);
let mw = BudgetMiddleware::new(BudgetLimits {
max_total_tokens: Some(100),
warn_fraction: Some(0.9),
..Default::default()
});
let tracker = mw.tracker();
let mut harness: AgentHarness<()> = AgentHarness::new();
harness
.register_model("mock", Arc::new(model))
.set_default_model("mock")
.register_tool(Arc::new(FakeTool::returning("noop", "ok")))
.push_middleware(Arc::new(mw));
let ctx = RunContext::new(RunConfig::new("budget-under"), ()).with_events(recorder.sink());
let run = harness
.invoke_in_context(&(), ctx, vec![Message::user("go")])
.await
.expect("a run under budget completes normally");
assert!(
run.final_response.is_some(),
"the run produced a final response"
);
assert_eq!(run.model_calls, 2, "two model turns completed");
assert_eq!(run.tool_calls, 1, "the noop tool ran once");
assert_eq!(tracker.snapshot().usage.usage.effective_total(), 10);
assert!(
!any_event(&recorder, |e| matches!(
e,
AgentEvent::BudgetExceeded { .. }
)),
"no BudgetExceeded event should be emitted under budget"
);
assert!(
!any_event(&recorder, |e| matches!(e, AgentEvent::BudgetWarning { .. })),
"no BudgetWarning event should be emitted well under the warn threshold"
);
}
#[tokio::test]
async fn input_reservation_blocks_oversized_call_in_a_live_run() {
let recorder = EventRecorder::new();
let mw = BudgetMiddleware::new(BudgetLimits {
max_input_tokens: Some(5),
..Default::default()
});
let mut harness: AgentHarness<()> = AgentHarness::new();
harness
.register_model("mock", Arc::new(MockModel::constant("unused")))
.set_default_model("mock")
.push_middleware(Arc::new(mw));
let big = "word ".repeat(200);
let ctx = RunContext::new(RunConfig::new("reserve-block"), ()).with_events(recorder.sink());
let err = harness
.invoke_in_context(&(), ctx, vec![Message::user(big)])
.await
.expect_err("the reservation preflight must block an oversized prompt");
assert!(
matches!(err, TinyAgentsError::LimitExceeded(_)),
"expected LimitExceeded, got {err:?}"
);
assert!(
any_event(&recorder, |e| matches!(
e,
AgentEvent::BudgetExceeded { blocked: true, .. }
)),
"the blocking reservation preflight must emit BudgetExceeded {{ blocked: true }}"
);
}
#[tokio::test]
async fn fitting_call_emits_reserved_and_reconciled_events() {
let recorder = EventRecorder::new();
let mw = BudgetMiddleware::new(BudgetLimits {
max_input_tokens: Some(1_000),
..Default::default()
});
let mut harness: AgentHarness<()> = AgentHarness::new();
harness
.register_model(
"mock",
Arc::new(MockModel::with_responses(vec![text_response("done", 3, 1)])),
)
.set_default_model("mock")
.push_middleware(Arc::new(mw));
let ctx = RunContext::new(RunConfig::new("reserve-ok"), ()).with_events(recorder.sink());
let run = harness
.invoke_in_context(&(), ctx, vec![Message::user("hi")])
.await
.expect("a small prompt fits the reservation and completes");
assert!(run.final_response.is_some(), "the run produced a response");
assert!(
any_event(&recorder, |e| matches!(
e,
AgentEvent::BudgetReserved { .. }
)),
"the preflight must emit BudgetReserved when max_input_tokens is set"
);
assert!(
any_event(&recorder, |e| matches!(
e,
AgentEvent::BudgetReconciled {
actual_input_tokens: 3,
..
}
)),
"after_model must reconcile the reservation against the actual 3 input tokens"
);
}
#[tokio::test]
async fn cached_input_budget_blocks_next_call() {
let recorder = EventRecorder::new();
let mut ctx = RunContext::new(RunConfig::new("cached-budget"), ()).with_events(recorder.sink());
let mw = BudgetMiddleware::new(BudgetLimits {
max_cached_input_tokens: Some(10),
..Default::default()
});
let tracker = mw.tracker();
let mut stack: MiddlewareStack<()> = MiddlewareStack::new();
stack.push(Arc::new(mw));
let mut resp = ModelResponse {
message: AssistantMessage {
id: None,
content: vec![ContentBlock::Text("cached".into())],
tool_calls: Vec::new(),
usage: Some(Usage {
cache_read_tokens: 12,
..Usage::new(2, 1)
}),
},
usage: Some(Usage {
cache_read_tokens: 12,
..Usage::new(2, 1)
}),
finish_reason: Some("stop".into()),
raw: None,
resolved_model: None,
};
stack
.run_after_model(&mut ctx, &(), &mut resp)
.await
.expect("recording cached usage does not fail");
assert_eq!(
tracker.snapshot().usage.usage.cache_read_tokens,
12,
"the tracker accumulates the reported cache-read tokens"
);
let mut req = ModelRequest::new(vec![Message::user("next")]);
let err = stack
.run_before_model(&mut ctx, &(), &mut req)
.await
.expect_err("the cached-input budget must block the next model call");
assert!(
matches!(err, TinyAgentsError::LimitExceeded(_)),
"expected LimitExceeded, got {err:?}"
);
assert!(
any_event(&recorder, |e| matches!(
e,
AgentEvent::BudgetExceeded { blocked: true, .. }
)),
"the blocking preflight emits BudgetExceeded {{ blocked: true }}"
);
}