use crate::error::{CredentialError, PromptError, ProviderError, ToolError};
use crate::ids::{ModelId, ProviderId};
use crate::policy::{
ApprovalDecision, ApprovalHandler, ApprovalRequest, Decision, Policy, PolicyRequest,
SteerSource,
};
use crate::prompt::{PromptContext, PromptFragment, PromptLayer};
use crate::provider::{
CredentialProvider, CredentialRequest, Credentials, ModelCapabilities, ModelProvider,
ModelRequest, ProviderEvent, ProviderStream, TokenCounter, TokenMeasurement,
TokenMeasurementSource,
};
use crate::tool::{
Concurrency, PreparedToolCall, Tool, ToolContext, ToolExecutionContext, ToolOutput, ToolSpec,
};
use async_trait::async_trait;
use futures::future::BoxFuture;
use futures::StreamExt;
use std::collections::VecDeque;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::{Arc, Mutex};
use tokio_util::sync::CancellationToken;
#[derive(Clone, Debug)]
pub enum ScriptedResponse {
Events(Vec<ProviderEvent>),
EventsThenHang(Vec<ProviderEvent>),
EventsThenError(Vec<ProviderEvent>, ProviderError),
Error(ProviderError),
Hang,
}
pub struct FakeProvider {
pub id: ProviderId,
pub capabilities: ModelCapabilities,
scripts: Mutex<VecDeque<ScriptedResponse>>,
requests: Mutex<Vec<ModelRequest>>,
}
impl FakeProvider {
pub fn new(capabilities: ModelCapabilities) -> Self {
Self {
id: ProviderId::from("fake"),
capabilities,
scripts: Mutex::new(VecDeque::new()),
requests: Mutex::new(Vec::new()),
}
}
pub fn push_events(&self, events: Vec<ProviderEvent>) {
self.scripts
.lock()
.unwrap()
.push_back(ScriptedResponse::Events(events));
}
pub fn push_events_then_hang(&self, events: Vec<ProviderEvent>) {
self.scripts
.lock()
.unwrap()
.push_back(ScriptedResponse::EventsThenHang(events));
}
pub fn push_events_then_error(&self, events: Vec<ProviderEvent>, error: ProviderError) {
self.scripts
.lock()
.unwrap()
.push_back(ScriptedResponse::EventsThenError(events, error));
}
pub fn push_error(&self, error: ProviderError) {
self.scripts
.lock()
.unwrap()
.push_back(ScriptedResponse::Error(error));
}
pub fn push_hang(&self) {
self.scripts
.lock()
.unwrap()
.push_back(ScriptedResponse::Hang);
}
pub fn requests(&self) -> Vec<ModelRequest> {
self.requests.lock().unwrap().clone()
}
}
#[async_trait]
impl ModelProvider for FakeProvider {
fn id(&self) -> ProviderId {
self.id.clone()
}
async fn capabilities(
&self,
_model: &ModelId,
_credentials: &dyn CredentialProvider,
_cancel: CancellationToken,
) -> Result<ModelCapabilities, ProviderError> {
Ok(self.capabilities)
}
async fn stream(
&self,
request: ModelRequest,
_credentials: &dyn CredentialProvider,
cancel: CancellationToken,
) -> Result<ProviderStream, ProviderError> {
let script = {
let mut scripts = self.scripts.lock().unwrap();
self.requests.lock().unwrap().push(request);
scripts.pop_front().unwrap_or_else(|| {
ScriptedResponse::Error(ProviderError::Unknown("no scripted response".into()))
})
};
let stream = match script {
ScriptedResponse::Events(events) => {
futures::stream::iter(events.into_iter().map(Ok).collect::<Vec<_>>()).boxed()
}
ScriptedResponse::EventsThenHang(events) => {
let cancel = cancel.clone();
futures::stream::iter(events.into_iter().map(Ok).collect::<Vec<_>>())
.chain(futures::stream::once(async move {
cancel.cancelled().await;
Err(ProviderError::Cancelled)
}))
.chain(futures::stream::pending())
.boxed()
}
ScriptedResponse::EventsThenError(events, error) => {
let mut items: Vec<Result<ProviderEvent, ProviderError>> =
events.into_iter().map(Ok).collect();
items.push(Err(error));
futures::stream::iter(items).boxed()
}
ScriptedResponse::Error(error) => futures::stream::iter(vec![Err(error)]).boxed(),
ScriptedResponse::Hang => {
let cancel = cancel.clone();
futures::stream::once(async move {
cancel.cancelled().await;
Err(ProviderError::Cancelled)
})
.chain(futures::stream::pending())
.boxed()
}
};
Ok(stream)
}
}
pub struct FakeTokenCounter {
pub estimator: Arc<dyn Fn(&ModelRequest) -> u64 + Send + Sync>,
pub source: TokenMeasurementSource,
}
impl FakeTokenCounter {
pub fn fixed(tokens: u64) -> Self {
Self {
estimator: Arc::new(move |_| tokens),
source: TokenMeasurementSource::Heuristic,
}
}
}
#[async_trait]
impl TokenCounter for FakeTokenCounter {
async fn count_input(
&self,
request: &ModelRequest,
_credentials: &dyn CredentialProvider,
_cancel: CancellationToken,
) -> Result<TokenMeasurement, ProviderError> {
Ok(TokenMeasurement {
input_tokens: (self.estimator)(request),
source: self.source,
safety_margin_tokens: 0,
})
}
}
#[derive(Default)]
pub struct FakeCredentialProvider {
pub resolve_calls: AtomicUsize,
}
#[async_trait]
impl CredentialProvider for FakeCredentialProvider {
async fn resolve(&self, _request: CredentialRequest) -> Result<Credentials, CredentialError> {
self.resolve_calls.fetch_add(1, Ordering::SeqCst);
Ok(Credentials::default())
}
}
type ToolBehavior = Arc<
dyn Fn(
PreparedToolCall,
ToolExecutionContext,
) -> BoxFuture<'static, Result<ToolOutput, ToolError>>
+ Send
+ Sync,
>;
pub struct ScriptedTool {
spec: ToolSpec,
behavior: ToolBehavior,
}
impl ScriptedTool {
pub fn echo(name: &str) -> Self {
Self::with_behavior(name, |call, _ctx| {
Box::pin(async move {
Ok(ToolOutput {
is_error: false,
text: call.arguments.to_string(),
})
})
})
}
pub fn failing(name: &str, message: &str) -> Self {
let message = message.to_string();
Self::with_behavior(name, move |_call, _ctx| {
let message = message.clone();
Box::pin(async move { Err(ToolError::Execution(message)) })
})
}
pub fn cancel_aware_sleep(name: &str, duration: std::time::Duration) -> Self {
Self::with_behavior(name, move |_call, ctx| {
let duration = duration;
Box::pin(async move {
tokio::select! {
_ = ctx.cancel.cancelled() => Err(ToolError::Cancelled),
_ = tokio::time::sleep(duration) => Ok(ToolOutput {
is_error: false,
text: "done".into(),
}),
}
})
})
}
pub fn parallel_sleepy(name: &str, duration: std::time::Duration) -> Self {
let mut tool = Self::cancel_aware_sleep(name, duration);
tool.spec.concurrency = Concurrency::ParallelSafe;
tool
}
pub fn with_behavior<F>(name: &str, behavior: F) -> Self
where
F: Fn(
PreparedToolCall,
ToolExecutionContext,
) -> BoxFuture<'static, Result<ToolOutput, ToolError>>
+ Send
+ Sync
+ 'static,
{
Self {
spec: ToolSpec {
name: name.into(),
description: "scripted test tool".into(),
parameters_schema: serde_json::json!({"type": "object"}),
concurrency: Concurrency::Sequential,
},
behavior: Arc::new(behavior),
}
}
}
#[async_trait]
impl Tool for ScriptedTool {
fn spec(&self) -> ToolSpec {
self.spec.clone()
}
async fn prepare(
&self,
arguments: serde_json::Value,
context: &ToolContext,
) -> Result<PreparedToolCall, ToolError> {
Ok(PreparedToolCall {
call_id: context.call_id.clone(),
name: self.spec.name.clone(),
arguments,
capabilities: Vec::new(),
})
}
async fn execute(
&self,
call: PreparedToolCall,
context: ToolExecutionContext,
) -> Result<ToolOutput, ToolError> {
(self.behavior)(call, context).await
}
}
pub struct ScriptedPolicy {
pub decisions: Mutex<VecDeque<Decision>>,
}
impl ScriptedPolicy {
pub fn allow() -> Self {
Self {
decisions: Mutex::new(VecDeque::new()),
}
}
pub fn deny(reason: &str) -> Self {
Self {
decisions: Mutex::new(VecDeque::from(vec![Decision::Deny {
reason: reason.into(),
}])),
}
}
pub fn ask() -> Self {
Self {
decisions: Mutex::new(VecDeque::from(vec![Decision::Ask])),
}
}
}
#[async_trait]
impl Policy for ScriptedPolicy {
async fn evaluate(
&self,
_request: PolicyRequest,
) -> Result<Decision, crate::error::PolicyError> {
Ok(self
.decisions
.lock()
.unwrap()
.pop_front()
.unwrap_or(Decision::Allow))
}
}
pub struct ScriptedApproval {
pub decision: ApprovalDecision,
pub wait_cancel_aware: bool,
}
#[async_trait]
impl ApprovalHandler for ScriptedApproval {
async fn wait(&self, _request: ApprovalRequest, cancel: CancellationToken) -> ApprovalDecision {
if self.wait_cancel_aware {
cancel.cancelled().await;
ApprovalDecision::Denied {
reason: "工具在执行前被取消".into(),
}
} else {
self.decision.clone()
}
}
}
pub struct StaticLayer {
pub id: String,
pub priority: i32,
pub content: Option<String>,
}
impl StaticLayer {
pub fn new(id: &str, priority: i32, content: Option<&str>) -> Self {
Self {
id: id.into(),
priority,
content: content.map(str::to_string),
}
}
}
#[async_trait]
impl PromptLayer for StaticLayer {
fn id(&self) -> &str {
&self.id
}
fn priority(&self) -> i32 {
self.priority
}
async fn render(
&self,
_context: &PromptContext,
) -> Result<Option<PromptFragment>, PromptError> {
Ok(self
.content
.clone()
.map(|content| PromptFragment { content }))
}
}
pub struct FailingLayer;
#[async_trait]
impl PromptLayer for FailingLayer {
fn id(&self) -> &str {
"failing"
}
fn priority(&self) -> i32 {
0
}
async fn render(
&self,
_context: &PromptContext,
) -> Result<Option<PromptFragment>, PromptError> {
Err(PromptError::Render("layer exploded".into()))
}
}
#[derive(Default)]
pub struct QueueSteerSource {
pub queue: Mutex<Vec<String>>,
}
impl QueueSteerSource {
pub fn new(items: Vec<String>) -> Self {
Self {
queue: Mutex::new(items),
}
}
}
#[async_trait]
impl SteerSource for QueueSteerSource {
async fn pending(&self) -> Vec<String> {
std::mem::take(self.queue.lock().unwrap().as_mut())
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::ids::{RunId, SessionId, ToolCallId, TurnId};
use crate::message::FinishReason;
use crate::provider::{ModelCapabilities, ModelRequest};
fn request() -> ModelRequest {
ModelRequest {
model: ModelId::from("test-model"),
system_prompt: String::new(),
messages: Vec::new(),
tools: Vec::new(),
reasoning: crate::provider::ReasoningLevel::Off,
generation: crate::provider::GenerationOptions::default(),
provider_options: serde_json::Value::Null,
}
}
#[tokio::test]
async fn fake_provider_replays_scripted_events_in_order() {
let provider = FakeProvider::new(ModelCapabilities {
context_tokens: 100_000,
max_output_tokens: 4_096,
supports_tools: true,
supports_images: true,
supports_reasoning: true,
});
provider.push_events(vec![
ProviderEvent::ResponseStarted,
ProviderEvent::TextDelta {
block: 0,
text: "hi".into(),
},
ProviderEvent::ResponseCompleted {
finish_reason: FinishReason::Stop,
},
]);
provider.push_error(ProviderError::Network("boom".into()));
let cancel = CancellationToken::new();
let creds = FakeCredentialProvider::default();
let mut stream = provider
.stream(request(), &creds, cancel.clone())
.await
.unwrap();
use futures::StreamExt;
let first = stream.next().await.unwrap().unwrap();
assert!(matches!(first, ProviderEvent::ResponseStarted));
let second = stream.next().await.unwrap().unwrap();
assert!(matches!(second, ProviderEvent::TextDelta { text, .. } if text == "hi"));
assert_eq!(provider.requests().len(), 1);
let mut second_stream = provider.stream(request(), &creds, cancel).await.unwrap();
assert!(matches!(
second_stream.next().await.unwrap().unwrap_err(),
ProviderError::Network(_)
));
}
#[tokio::test]
async fn hang_stream_yields_cancelled_error_when_cancelled() {
let provider = FakeProvider::new(ModelCapabilities {
context_tokens: 100_000,
max_output_tokens: 4_096,
supports_tools: true,
supports_images: true,
supports_reasoning: true,
});
provider.push_hang();
let cancel = CancellationToken::new();
let creds = FakeCredentialProvider::default();
let mut stream = provider
.stream(request(), &creds, cancel.clone())
.await
.unwrap();
cancel.cancel();
use futures::StreamExt;
let item = stream.next().await.unwrap();
assert!(matches!(item, Err(ProviderError::Cancelled)));
}
#[tokio::test]
async fn fake_token_counter_uses_estimator() {
let counter = FakeTokenCounter::fixed(42);
let cancel = CancellationToken::new();
let creds = FakeCredentialProvider::default();
let m = counter
.count_input(&request(), &creds, cancel)
.await
.unwrap();
assert_eq!(m.input_tokens, 42);
assert_eq!(m.source, TokenMeasurementSource::Heuristic);
}
#[tokio::test]
async fn echo_tool_returns_arguments() {
let tool = ScriptedTool::echo("echo");
let ctx = ToolContext {
session_id: SessionId::from("s"),
run_id: RunId::from("r"),
turn_id: TurnId::from("t"),
call_id: ToolCallId::from("c1"),
};
let prepared = tool
.prepare(serde_json::json!({"a": 1}), &ctx)
.await
.unwrap();
assert_eq!(prepared.call_id, ToolCallId::from("c1"));
let (tx, _rx) = tokio::sync::mpsc::channel(4);
let out = tool
.execute(
prepared,
ToolExecutionContext {
cancel: CancellationToken::new(),
progress: tx,
},
)
.await
.unwrap();
assert_eq!(out.text, "{\"a\":1}");
}
#[tokio::test]
async fn cancel_aware_tool_returns_cancelled_error() {
let tool = ScriptedTool::cancel_aware_sleep("slow", std::time::Duration::from_secs(30));
let ctx = ToolContext {
session_id: SessionId::from("s"),
run_id: RunId::from("r"),
turn_id: TurnId::from("t"),
call_id: ToolCallId::from("c1"),
};
let prepared = tool.prepare(serde_json::json!({}), &ctx).await.unwrap();
let cancel = CancellationToken::new();
cancel.cancel();
let (tx, _rx) = tokio::sync::mpsc::channel(4);
let result = tool
.execute(
prepared,
ToolExecutionContext {
cancel,
progress: tx,
},
)
.await;
assert!(matches!(result, Err(ToolError::Cancelled)));
}
}