use crate::api::ApiClient;
use crate::api::error::ApiError;
use crate::config::SessionConfig;
use crate::message::{Message, MessagePart, Role};
use crate::stream::{
DeltaPart, IndexedDelta, MessageDelta, MessageDeltaPayload, MessageMetadata, MessageStart,
PartStart, StreamEvent, Usage,
};
use crate::tool::{Tool, ToolContext, ToolError, ToolOutput, ToolSchema};
use futures::Stream;
use serde_json::{Value, json};
use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use std::sync::Mutex;
#[derive(Clone)]
pub struct MockApiClient {
model_name: Arc<std::sync::Mutex<String>>,
responses: Arc<Mutex<Vec<MockResponse>>>,
error: Option<String>,
}
#[derive(Clone, Default)]
pub struct MockResponse {
pub text: String,
pub tool_call: Option<MockToolCall>,
pub stop_reason: String,
}
#[derive(Clone)]
pub struct MockToolCall {
pub id: String,
pub name: String,
pub input: Value,
}
impl MockApiClient {
#[must_use]
pub fn new(model: &str) -> Self {
let default_response = MockResponse {
text: "Hello!".to_string(),
tool_call: None,
stop_reason: "end_turn".to_string(),
};
Self {
model_name: Arc::new(std::sync::Mutex::new(model.to_string())),
responses: Arc::new(Mutex::new(vec![default_response])),
error: None,
}
}
#[must_use]
pub fn with_text_response(self, text: &str) -> Self {
if let Some(r) = crate::error::recover_guard(self.responses.lock()).first_mut() {
r.text = text.to_string();
}
self
}
#[must_use]
pub fn with_tool_call(self, id: &str, name: &str, input: Value) -> Self {
let mut responses = crate::error::recover_guard(self.responses.lock());
if let Some(r) = responses.first_mut() {
r.tool_call = Some(MockToolCall {
id: id.to_string(),
name: name.to_string(),
input,
});
r.stop_reason = "tool_use".to_string();
}
drop(responses);
self
}
#[must_use]
pub fn with_stop_reason(self, reason: &str) -> Self {
if let Some(r) = crate::error::recover_guard(self.responses.lock()).first_mut() {
r.stop_reason = reason.to_string();
}
self
}
#[must_use]
pub fn with_responses(self, responses: Vec<MockResponse>) -> Self {
if !responses.is_empty() {
*crate::error::recover_guard(self.responses.lock()) = responses;
}
self
}
#[must_use]
pub fn with_error(mut self, error: &str) -> Self {
self.error = Some(error.to_string());
self
}
fn stream_events(&self, model: String) -> Vec<Result<StreamEvent, ApiError>> {
if let Some(ref err) = self.error {
return vec![Err(ApiError::api(err))];
}
let response = self.pop_response();
let mut events: Vec<Result<StreamEvent, ApiError>> =
vec![Ok(StreamEvent::MessageStart(MessageStart {
message: MessageMetadata {
id: "msg_test".to_string(),
role: "assistant".to_string(),
model,
},
}))];
let text = response.text.clone();
events.push(Ok(StreamEvent::PartStart(PartStart {
index: 0,
part: Some(MessagePart::text("")),
})));
events.push(Ok(StreamEvent::IndexedDelta(IndexedDelta {
index: 0,
delta: DeltaPart::Text { text },
})));
events.push(Ok(StreamEvent::PartStop { index: Some(0) }));
if let Some(tc) = &response.tool_call {
events.push(Ok(StreamEvent::PartStart(PartStart {
index: 1,
part: Some(MessagePart::tool_call(
&tc.id,
&tc.name,
serde_json::json!({}),
)),
})));
events.push(Ok(StreamEvent::IndexedDelta(IndexedDelta {
index: 1,
delta: DeltaPart::InputJson {
partial_json: tc.input.to_string(),
},
})));
events.push(Ok(StreamEvent::PartStop { index: Some(1) }));
}
events.push(Ok(StreamEvent::MessageDelta(MessageDelta {
delta: MessageDeltaPayload {
stop_reason: Some(response.stop_reason),
},
usage: Some(Usage::new(50, 25)),
})));
events.push(Ok(StreamEvent::MessageStop));
events
}
fn pop_response(&self) -> MockResponse {
let mut guard = crate::error::recover_guard(self.responses.lock());
if guard.len() > 1 {
guard.remove(0)
} else {
guard.first().cloned().unwrap_or_default()
}
}
}
fn unsupported_mock_option(options: &crate::structured::RequestOptions) -> Option<ApiError> {
if !matches!(
options.tool_constraint,
crate::structured::ToolConstraint::None
) {
return Some(ApiError::config(
"this client does not support tool-call constraints (tool_constraint)",
));
}
None
}
impl ApiClient for MockApiClient {
fn model(&self) -> String {
crate::error::recover_guard(self.model_name.lock()).clone()
}
fn set_model(&self, model: &str) -> bool {
if model.trim().is_empty() {
return false;
}
*crate::error::recover_guard(self.model_name.lock()) = model.to_string();
true
}
fn stream_messages(
&self,
_request: &crate::api::StreamRequest,
) -> Pin<Box<dyn Stream<Item = Result<StreamEvent, ApiError>> + Send + 'static>> {
let model = crate::error::recover_guard(self.model_name.lock()).clone();
let events = self.stream_events(model);
Box::pin(futures::stream::iter(events))
}
fn stream_messages_with_options(
&self,
_request: &crate::api::StreamRequest,
options: crate::structured::RequestOptions,
) -> Pin<Box<dyn Stream<Item = Result<StreamEvent, ApiError>> + Send + 'static>> {
if let Some(err) = unsupported_mock_option(&options) {
return Box::pin(futures::stream::once(async move { Err(err) }));
}
let model = options
.model
.unwrap_or_else(|| crate::error::recover_guard(self.model_name.lock()).clone());
let events = self.stream_events(model);
Box::pin(futures::stream::iter(events))
}
fn create_message(
&self,
_request: &crate::api::StreamRequest,
) -> Pin<Box<dyn Future<Output = Result<crate::api::NonStreamingResponse, ApiError>> + Send + '_>>
{
if let Some(ref err) = self.error {
let err = err.clone();
return Box::pin(async move { Err(ApiError::api(&err)) });
}
let response = self.pop_response();
Box::pin(async move {
let stop_reason = crate::stream::StreamStopReason::from_api_str(&response.stop_reason)
.unwrap_or(crate::stream::StreamStopReason::EndTurn);
let mut parts = vec![MessagePart::text(response.text)];
if let Some(tc) = response.tool_call {
parts.push(MessagePart::tool_call(tc.id, tc.name, tc.input));
}
Ok(crate::api::NonStreamingResponse {
message: Message::new(Role::Assistant, parts),
stop_reason,
usage: Some(crate::stream::Usage::new(50, 25)),
})
})
}
fn create_message_with_options(
&self,
request: &crate::api::StreamRequest,
options: crate::structured::RequestOptions,
) -> Pin<Box<dyn Future<Output = Result<crate::api::NonStreamingResponse, ApiError>> + Send + '_>>
{
if let Some(err) = unsupported_mock_option(&options) {
return Box::pin(async move { Err(err) });
}
self.create_message(request)
}
}
pub struct MockTool {
name: String,
description: String,
input_schema: Value,
result: String,
is_error: bool,
is_concurrency_safe: bool,
is_read_only: bool,
delay: std::time::Duration,
system_prompt: Option<String>,
}
impl MockTool {
#[must_use]
pub fn new(name: &str, description: &str) -> Self {
Self {
name: name.to_string(),
description: description.to_string(),
input_schema: json!({
"type": "object",
"properties": { "input": { "type": "string" } }
}),
result: "mock result".to_string(),
is_error: false,
is_concurrency_safe: false,
is_read_only: true,
delay: std::time::Duration::ZERO,
system_prompt: None,
}
}
#[must_use]
pub fn with_result(mut self, result: &str) -> Self {
self.result = result.to_string();
self
}
#[must_use]
pub fn with_error(mut self) -> Self {
self.is_error = true;
self
}
#[must_use]
pub fn with_concurrency_safe(mut self, safe: bool) -> Self {
self.is_concurrency_safe = safe;
self
}
#[must_use]
pub fn with_read_only(mut self, read_only: bool) -> Self {
self.is_read_only = read_only;
self
}
#[must_use]
pub fn with_delay(mut self, delay: std::time::Duration) -> Self {
self.delay = delay;
self
}
#[must_use]
pub fn with_schema(mut self, schema: Value) -> Self {
self.input_schema = schema;
self
}
#[must_use]
pub fn with_system_prompt(mut self, prompt: &str) -> Self {
self.system_prompt = Some(prompt.to_string());
self
}
}
impl Tool for MockTool {
fn name(&self) -> &str {
&self.name
}
fn description(&self) -> &str {
&self.description
}
fn schema(&self) -> ToolSchema {
ToolSchema {
tool: self.name.clone(),
description: self.description.clone(),
input_schema: self.input_schema.clone(),
}
}
fn call(
&self,
_input: Value,
_context: &ToolContext,
) -> Pin<Box<dyn Future<Output = Result<ToolOutput, ToolError>> + Send + '_>> {
let result = self.result.clone();
let is_error = self.is_error;
let delay = self.delay;
Box::pin(async move {
if !delay.is_zero() {
tokio::time::sleep(delay).await;
}
if is_error {
Err(ToolError::Execution(result))
} else {
Ok(ToolOutput::text(result))
}
})
}
fn is_concurrency_safe(&self) -> bool {
self.is_concurrency_safe
}
fn is_read_only(&self) -> bool {
self.is_read_only
}
fn system_prompt(&self) -> Option<String> {
self.system_prompt.clone()
}
}
#[must_use]
pub fn test_message(text: &str) -> Message {
Message::user(text)
}
#[must_use]
pub fn test_assistant_message(text: &str) -> Message {
Message::assistant(text)
}
#[must_use]
pub fn test_tool_use_message(calls: &[(&str, &str, Value)]) -> Message {
let blocks: Vec<MessagePart> = calls
.iter()
.enumerate()
.map(|(i, (id, name, input))| {
MessagePart::tool_call(
if id.is_empty() {
format!("call_{i}")
} else {
id.to_string()
},
*name,
input.clone(),
)
})
.collect();
Message::new(Role::Assistant, blocks)
}
#[must_use]
pub fn test_config() -> SessionConfig {
SessionConfig {
system_prompt: Some("You are a test assistant.".to_string()),
..SessionConfig::default()
}
}
#[cfg(test)]
mod tests {
#[derive(Debug, serde::Deserialize, PartialEq)]
struct Forecast {
role: String,
}
impl crate::structured::StructuredOutput for Forecast {
fn name() -> &'static str {
"forecast"
}
fn schema() -> serde_json::Value {
serde_json::json!({
"type": "object",
"properties": {"role": {"type": "string"}},
"required": ["role"]
})
}
}
#[tokio::test]
async fn request_structured_against_the_mock_returns_the_deserialized_type() {
let client = MockApiClient::new("test-model").with_text_response(r#"{"role":"coding"}"#);
let got = crate::structured::request_structured::<Forecast>(
&client,
vec![crate::message::Message::user("hi")],
None,
)
.await;
match got {
Ok(forecast) => assert_eq!(forecast.role, "coding"),
Err(e) => panic!(
"the mock must serve the canned text to a structured \
request: {e:?}"
),
}
}
#[tokio::test]
async fn request_structured_against_a_failing_mock_returns_the_api_error() {
let client = MockApiClient::new("test-model").with_error("boom");
let got = crate::structured::request_structured::<Forecast>(
&client,
vec![crate::message::Message::user("hi")],
None,
)
.await;
assert!(
matches!(got, Err(crate::structured::StructuredError::Api(_))),
"a failing provider surfaces as StructuredError::Api, got {got:?}"
);
}
#[tokio::test]
async fn request_structured_with_malformed_json_returns_deserialize() {
let client = MockApiClient::new("test-model").with_text_response("not json at all");
let got = crate::structured::request_structured::<Forecast>(
&client,
vec![crate::message::Message::user("hi")],
None,
)
.await;
assert!(
matches!(got, Err(crate::structured::StructuredError::Deserialize(_))),
"a struct type cannot fall back to prose, so malformed JSON \
reaches the caller as Deserialize, got {got:?}"
);
}
#[tokio::test]
async fn streaming_serves_a_structured_request_instead_of_rejecting_it() {
use futures::StreamExt;
let client = MockApiClient::new("test-model").with_text_response("hi");
let opts = crate::structured::RequestOptions::default()
.with_response_format(crate::structured::ResponseFormat::from_type::<Forecast>());
let stream =
client.stream_messages_with_options(&crate::api::StreamRequest::new(vec![]), opts);
let events: Vec<_> = stream.collect().await;
assert!(
!events.is_empty() && events.iter().all(Result::is_ok),
"the canned stream is served under a response_format, got {events:?}"
);
}
#[tokio::test]
async fn the_mock_rejects_a_tool_constraint_loudly() {
let client = MockApiClient::new("test-model").with_text_response("hi");
let opts = crate::structured::RequestOptions {
tool_constraint: crate::structured::ToolConstraint::Strict,
..Default::default()
};
let err = client
.create_message_with_options(&crate::api::StreamRequest::new(vec![]), opts)
.await
.expect_err("a tool constraint must be rejected loudly");
assert!(
err.to_string().contains("tool_constraint"),
"the error names the option: {err}"
);
}
#[tokio::test]
async fn streaming_serves_a_structured_request_under_a_model_override() {
use futures::StreamExt;
let client = MockApiClient::new("test-model").with_text_response("hi");
let opts = crate::structured::RequestOptions::default()
.with_model("fallback")
.with_response_format(crate::structured::ResponseFormat::from_type::<Forecast>());
let stream =
client.stream_messages_with_options(&crate::api::StreamRequest::new(vec![]), opts);
let events: Vec<_> = stream.collect().await;
let model = events.iter().find_map(|e| {
e.as_ref().ok().and_then(|ev| match ev {
crate::stream::StreamEvent::MessageStart(start) => {
Some(start.message.model.clone())
}
_ => None,
})
});
assert_eq!(
model.as_deref(),
Some("fallback"),
"accepting a response_format must not disturb the model \
override routing: {events:?}"
);
}
#[tokio::test]
async fn streaming_still_rejects_a_tool_constraint() {
use futures::StreamExt;
let client = MockApiClient::new("test-model").with_text_response("hi");
let opts = crate::structured::RequestOptions {
tool_constraint: crate::structured::ToolConstraint::Strict,
..Default::default()
};
let stream =
client.stream_messages_with_options(&crate::api::StreamRequest::new(vec![]), opts);
let events: Vec<_> = stream.collect().await;
let err = events
.first()
.and_then(|e| e.as_ref().err())
.unwrap_or_else(|| {
panic!("the streaming path must reject a constraint, got {events:?}")
});
assert!(
err.to_string().contains("tool_constraint"),
"the error names the option: {err}"
);
}
use super::*;
use crate::tool::ToolRegistry;
use futures::StreamExt;
#[test]
fn test_mock_client_model() {
let client = MockApiClient::new("test-model");
assert_eq!(client.model(), "test-model");
}
#[tokio::test]
async fn test_mock_client_default_response() {
let client = MockApiClient::new("test-model");
let stream = client.stream_messages(&crate::api::StreamRequest {
messages: vec![Message::user("Hi")],
system: None,
tools: None,
});
let events: Vec<_> = stream.collect().await;
assert!(events.len() >= 4);
assert!(events[0].is_ok());
}
#[tokio::test]
async fn test_mock_client_custom_text() {
let client = MockApiClient::new("test-model").with_text_response("Custom response");
let stream = client.stream_messages(&crate::api::StreamRequest {
messages: vec![Message::user("Hi")],
system: None,
tools: None,
});
let events: Vec<_> = stream.collect().await;
let has_text = events.iter().any(|e| {
if let Ok(StreamEvent::IndexedDelta(delta)) = e
&& let DeltaPart::Text { text } = &delta.delta
{
return text == "Custom response";
}
false
});
assert!(has_text);
}
#[tokio::test]
async fn test_mock_client_tool_call() {
let client = MockApiClient::new("test-model").with_tool_call(
"call_1",
"echo",
json!({"message": "hi"}),
);
let stream = client.stream_messages(&crate::api::StreamRequest {
messages: vec![Message::user("Hi")],
system: None,
tools: None,
});
let events: Vec<_> = stream.collect().await;
let has_tool_use = events.iter().any(|e| {
if let Ok(StreamEvent::PartStart(start)) = e
&& let Some(MessagePart::ToolCall { name, .. }) = &start.part
{
return name == "echo";
}
false
});
assert!(has_tool_use);
let has_tool_stop = events.iter().any(|e| {
if let Ok(StreamEvent::MessageDelta(delta)) = e {
delta.delta.stop_reason.as_deref() == Some("tool_use")
} else {
false
}
});
assert!(has_tool_stop);
let mut accumulator = crate::stream::StreamAccumulator::new();
for event in &events {
let Ok(event) = event else {
panic!("mock stream must not carry errors here");
};
accumulator
.process(event)
.expect("accumulator accepts mock events");
}
let assembled = accumulator.build();
let tool_calls = assembled.tool_call_parts();
assert_eq!(
tool_calls.len(),
1,
"one tool call reconstructed from the stream"
);
let (id, name, input) = tool_calls[0];
assert_eq!(id, "call_1");
assert_eq!(name, "echo");
assert_eq!(input, &json!({"message": "hi"}));
}
#[tokio::test]
async fn test_mock_client_error() {
let client = MockApiClient::new("test-model").with_error("API error");
let stream = client.stream_messages(&crate::api::StreamRequest {
messages: vec![Message::user("Hi")],
system: None,
tools: None,
});
let events: Vec<_> = stream.collect().await;
assert_eq!(events.len(), 1);
assert!(events[0].is_err());
}
#[tokio::test]
async fn test_mock_client_multi_turn() {
let client = MockApiClient::new("test-model").with_responses(vec![
MockResponse {
text: "First".to_string(),
tool_call: None,
stop_reason: "end_turn".to_string(),
},
MockResponse {
text: "Second".to_string(),
tool_call: None,
stop_reason: "end_turn".to_string(),
},
]);
let stream1 = client.stream_messages(&crate::api::StreamRequest {
messages: vec![Message::user("Hi")],
system: None,
tools: None,
});
let events1: Vec<_> = stream1.collect().await;
let has_first = events1.iter().any(|e| {
if let Ok(StreamEvent::IndexedDelta(delta)) = e
&& let DeltaPart::Text { text } = &delta.delta
{
return text == "First";
}
false
});
assert!(has_first);
let stream2 = client.stream_messages(&crate::api::StreamRequest {
messages: vec![Message::user("Hi")],
system: None,
tools: None,
});
let events2: Vec<_> = stream2.collect().await;
let has_second = events2.iter().any(|e| {
if let Ok(StreamEvent::IndexedDelta(delta)) = e
&& let DeltaPart::Text { text } = &delta.delta
{
return text == "Second";
}
false
});
assert!(has_second);
}
#[tokio::test]
async fn test_mock_client_create_message() {
let client = MockApiClient::new("test-model").with_text_response("Hello!");
let result = client
.create_message(&crate::api::StreamRequest {
messages: vec![Message::user("Hi")],
system: None,
tools: None,
})
.await;
assert!(result.is_ok());
let response = result.unwrap();
assert_eq!(response.message.text_content(), "Hello!");
assert_eq!(
response.stop_reason,
crate::stream::StreamStopReason::EndTurn
);
}
#[tokio::test]
async fn create_message_keeps_text_alongside_tool_call() {
let client = MockApiClient::new("test-model").with_responses(vec![MockResponse {
text: "Let me look that up.".to_string(),
tool_call: Some(MockToolCall {
id: "call_1".to_string(),
name: "search".to_string(),
input: json!({}),
}),
stop_reason: "tool_use".to_string(),
}]);
let response = client
.create_message(&crate::api::StreamRequest::new(vec![]))
.await
.expect("mock must respond");
assert_eq!(
response.message.text_content(),
"Let me look that up.",
"doc: the assistant message is built from the response's text and optional tool call"
);
assert!(
response.message.tool_call_parts().len() == 1,
"the tool call rides alongside the text, mirroring the streaming twin"
);
}
#[tokio::test]
async fn mock_usage_matches_across_streaming_and_non_streaming() {
use futures::StreamExt;
let stream_client = MockApiClient::new("test-model").with_text_response("hi");
let mut stream = stream_client.stream_messages(&crate::api::StreamRequest::new(vec![]));
let mut streamed_usage = None;
while let Some(item) = stream.next().await {
if let Ok(crate::stream::StreamEvent::MessageDelta(delta)) = item {
streamed_usage = delta.usage;
}
}
let create_client = MockApiClient::new("test-model").with_text_response("hi");
let response = create_client
.create_message(&crate::api::StreamRequest::new(vec![]))
.await
.expect("mock must respond");
assert_eq!(
streamed_usage, response.usage,
"one mock, one usage story: the streaming twin reports \
{streamed_usage:?} while create_message reports {:?}",
response.usage
);
}
#[tokio::test]
async fn one_mock_response_builds_the_same_message_on_both_paths() {
use futures::StreamExt;
let scripted = MockResponse {
text: "Let me look that up.".to_string(),
tool_call: Some(MockToolCall {
id: "call_1".to_string(),
name: "search".to_string(),
input: json!({"query": "loop"}),
}),
stop_reason: "tool_use".to_string(),
};
let client = MockApiClient::new("test-model").with_responses(vec![scripted]);
let mut stream = client.stream_messages(&crate::api::StreamRequest::new(vec![]));
let mut accumulator = crate::stream::StreamAccumulator::new();
while let Some(item) = stream.next().await {
accumulator
.process(&item.expect("mock stream must not carry errors"))
.expect("accumulator accepts mock events");
}
let streamed = accumulator.build();
let response = client
.create_message(&crate::api::StreamRequest::new(vec![]))
.await
.expect("mock must respond");
assert_eq!(
streamed.text_content(),
response.message.text_content(),
"the same scripted response must produce the same text on both paths"
);
assert_eq!(
streamed.tool_call_parts(),
response.message.tool_call_parts(),
"the same scripted response must produce the same tool call on both paths"
);
}
#[tokio::test]
async fn test_mock_tool() {
let tool = MockTool::new("echo", "Echoes input").with_result("Echo: hello");
let ctx = ToolContext::default();
let result = tool.call(json!({"input": "hello"}), &ctx).await;
assert!(result.is_ok());
assert_eq!(result.unwrap().text_content(), "Echo: hello");
}
#[tokio::test]
async fn test_mock_tool_error() {
let tool = MockTool::new("fail", "Always fails")
.with_result("Something went wrong")
.with_error();
let ctx = ToolContext::default();
let result = tool.call(json!({}), &ctx).await;
assert!(result.is_err());
}
#[test]
fn test_mock_tool_registry() {
let tool = MockTool::new("echo", "Echoes input").with_concurrency_safe(true);
let mut registry = ToolRegistry::new();
registry.register(tool);
assert!(registry.contains("echo"));
assert_eq!(registry.len(), 1);
let t = registry.get("echo").unwrap();
assert_eq!(t.name(), "echo");
assert!(t.is_concurrency_safe());
assert!(t.is_read_only());
}
#[test]
fn test_fixture_test_message() {
let msg = test_message("Hello");
assert_eq!(msg.role, Role::User);
assert_eq!(msg.parts.len(), 1);
}
#[test]
fn test_fixture_test_assistant_message() {
let msg = test_assistant_message("Hi there");
assert_eq!(msg.role, Role::Assistant);
}
#[test]
fn test_fixture_test_tool_use_message() {
let msg = test_tool_use_message(&[("call_1", "bash", json!({"command": "ls"}))]);
assert_eq!(msg.role, Role::Assistant);
assert_eq!(msg.parts.len(), 1);
assert!(msg.parts[0].is_tool_call());
}
#[test]
fn test_fixture_test_config() {
let config = test_config();
assert_eq!(
config.system_prompt.as_deref(),
Some("You are a test assistant.")
);
}
#[test]
fn mock_api_client_set_model() {
let client = MockApiClient::new("model-a");
assert_eq!(client.model(), "model-a");
assert!(client.set_model("model-b"));
assert_eq!(client.model(), "model-b");
}
#[test]
fn mock_api_client_set_model_rejects_empty() {
let client = MockApiClient::new("model-a");
assert!(!client.set_model(""));
assert!(!client.set_model(" "));
assert_eq!(client.model(), "model-a");
}
}
static ENV_LOCK: std::sync::Mutex<()> = std::sync::Mutex::new(());
thread_local! {
static ENV_GUARD_HELD: std::cell::Cell<bool> = const { std::cell::Cell::new(false) };
}
pub struct EnvGuard {
_lock: std::sync::MutexGuard<'static, ()>,
saved: Vec<(&'static str, Option<std::ffi::OsString>)>,
}
impl EnvGuard {
pub fn acquire(vars: &[&'static str]) -> Self {
ENV_GUARD_HELD.with(|held| {
assert!(
!held.get(),
"nested EnvGuard::acquire would deadlock; one guard per test"
);
});
let lock = ENV_LOCK
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
let saved = vars
.iter()
.map(|name| (*name, std::env::var_os(name)))
.collect();
ENV_GUARD_HELD.with(|held| held.set(true));
Self { _lock: lock, saved }
}
pub fn set(&self, name: &'static str, value: &str) {
assert!(
self.saved
.iter()
.any(|(snapshotted, _)| *snapshotted == name),
"setting {name} without a snapshot: Drop would not restore it"
);
unsafe { std::env::set_var(name, value) }
}
pub fn remove(&self, name: &'static str) {
assert!(
self.saved
.iter()
.any(|(snapshotted, _)| *snapshotted == name),
"removing {name} without a snapshot: Drop would not restore it"
);
unsafe { std::env::remove_var(name) }
}
}
impl Drop for EnvGuard {
fn drop(&mut self) {
ENV_GUARD_HELD.with(|held| held.set(false));
for (name, value) in &self.saved {
match value {
Some(previous) => unsafe { std::env::set_var(name, previous) },
None => unsafe { std::env::remove_var(name) },
}
}
}
}
#[cfg(test)]
mod env_guard_tests {
use super::*;
#[test]
fn guard_restores_previous_values_on_normal_exit() {
{
let _lock = ENV_LOCK
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
unsafe { std::env::set_var("LOOPCTL_GUARD_VALUE", "original") };
}
{
let env = EnvGuard::acquire(&["LOOPCTL_GUARD_VALUE"]);
env.set("LOOPCTL_GUARD_VALUE", "mutated");
env.remove("LOOPCTL_GUARD_VALUE");
env.set("LOOPCTL_GUARD_VALUE", "mutated-again");
}
let restored = {
let _lock = ENV_LOCK
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
std::env::var("LOOPCTL_GUARD_VALUE")
};
assert_eq!(
restored,
Ok("original".to_string()),
"drop must restore the pre-guard value after mixed set/remove churn"
);
{
let _lock = ENV_LOCK
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
unsafe { std::env::remove_var("LOOPCTL_GUARD_VALUE") };
}
}
#[test]
fn nested_acquire_fails_fast() {
let result = std::panic::catch_unwind(|| {
let _outer = EnvGuard::acquire(&["LOOPCTL_GUARD_NESTED"]);
let _inner = EnvGuard::acquire(&["LOOPCTL_GUARD_NESTED"]);
});
let panic = result.expect_err("nested acquire must panic, not deadlock");
let names_it = panic
.downcast_ref::<String>()
.is_some_and(|m| m.contains("one guard per test"))
|| panic
.downcast_ref::<&str>()
.is_some_and(|m| m.contains("one guard per test"));
assert!(
names_it,
"the panic must name the nesting rule so the failure is actionable"
);
}
#[test]
fn mutating_unsnapshotted_variable_fails() {
let result = std::panic::catch_unwind(|| {
let env = EnvGuard::acquire(&["LOOPCTL_GUARD_SNAPSHOT"]);
env.set("LOOPCTL_GUARD_UNSNAPSHOTTED", "x");
});
let panic = result.expect_err("mutating an unsnapshotted name must panic");
let names_it = panic
.downcast_ref::<String>()
.is_some_and(|m| m.contains("without a snapshot"))
|| panic
.downcast_ref::<&str>()
.is_some_and(|m| m.contains("without a snapshot"));
assert!(
names_it,
"the panic must name the snapshot rule so the failure is actionable"
);
}
#[test]
fn guard_restores_after_panic() {
{
let env = EnvGuard::acquire(&["LOOPCTL_GUARD_TEST"]);
env.remove("LOOPCTL_GUARD_TEST");
}
let result = std::panic::catch_unwind(|| {
let env = EnvGuard::acquire(&["LOOPCTL_GUARD_TEST"]);
env.set("LOOPCTL_GUARD_TEST", "mutated");
panic!("simulated test failure");
});
assert!(result.is_err(), "the closure must have panicked");
assert!(
std::env::var("LOOPCTL_GUARD_TEST").is_err(),
"drop during unwinding must restore the variable to absent"
);
}
}