Files
cassady/tests/embedding_tests.rs
T
owen 71f84c03ce
CI / Build (push) Waiting to run
CI / Test (push) Waiting to run
Add Rust embedding API
2026-06-25 06:53:13 -05:00

316 lines
9.9 KiB
Rust

use cassady::access::AccessMode;
use cassady::config::ReasoningEffort;
use cassady::conversation::Record;
use cassady::embedding::{Event, SessionBuilder};
use serde_json::json;
use tempfile::tempdir;
use wiremock::matchers::{body_string_contains, method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
fn sse(body: &str) -> ResponseTemplate {
ResponseTemplate::new(200).set_body_raw(body.as_bytes().to_vec(), "text/event-stream")
}
fn content_sse(content: &str) -> ResponseTemplate {
sse(&format!(
"data: {{\"choices\":[{{\"index\":0,\"delta\":{{\"content\":{}}}}}]}}\r\n\r\ndata: [DONE]\r\n\r\n",
serde_json::to_string(content).unwrap()
))
}
fn tool_call_sse(id: &str, name: &str, arguments: &str) -> ResponseTemplate {
sse(&format!(
"data: {{\"choices\":[{{\"index\":0,\"delta\":{{\"tool_calls\":[{{\"index\":0,\"id\":\"{id}\",\"type\":\"function\",\"function\":{{\"name\":\"{name}\",\"arguments\":{}}}}}]}}}}]}}\r\n\r\ndata: [DONE]\r\n\r\n",
serde_json::to_string(arguments).unwrap()
))
}
fn write_test_config(root: &std::path::Path, base_url: &str) {
std::fs::write(
root.join("providers.json"),
serde_json::to_string_pretty(&json!({
"providers": [{
"id": "test-provider",
"kind": "openai-compatible",
"base_url": base_url,
"api_key": "test-key",
"default_model": "test-model",
"models": ["test-model"]
}]
}))
.unwrap(),
)
.unwrap();
std::fs::write(
root.join("models.json"),
serde_json::to_string_pretty(&json!({
"models": [{
"id": "test-model",
"provider": "test-provider",
"context_length": 128,
"max_output_tokens": 64,
"reasoning": {
"supported": true,
"required": false,
"default_effort": "off",
"request_format": "reasoning_effort"
}
}]
}))
.unwrap(),
)
.unwrap();
std::fs::write(
root.join("config.json"),
serde_json::to_string_pretty(&json!({
"default_provider": "test-provider",
"default_model": "test-model",
"default_reasoning_effort": "off"
}))
.unwrap(),
)
.unwrap();
}
#[tokio::test]
async fn embedded_session_runs_turn_and_streams_events() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/chat/completions"))
.respond_with(content_sse("Hello from embedded Cassady."))
.expect(1)
.mount(&server)
.await;
let root = tempdir().unwrap();
let cwd = tempdir().unwrap();
write_test_config(root.path(), &server.uri());
let session = SessionBuilder::new()
.config_root(root.path())
.cwd(cwd.path())
.access_mode(AccessMode::ReadOnly)
.reasoning_effort(ReasoningEffort::Off)
.build()
.await
.unwrap();
assert_eq!(session.model(), "test-model");
assert_eq!(session.access_mode(), AccessMode::ReadOnly);
let mut turn = session.start_turn("say hi").await.unwrap();
let mut streamed = String::new();
while let Some(event) = turn.next_event().await.unwrap() {
match event {
Event::AssistantChunk(chunk) => streamed.push_str(&chunk),
Event::Finished => break,
_ => {}
}
}
let session = turn.finish().await.unwrap();
assert_eq!(streamed, "Hello from embedded Cassady.");
assert!(session.records().iter().any(|record| matches!(
record,
Record::Assistant { content, .. } if content == "Hello from embedded Cassady."
)));
assert!(session.conversation_path().is_file());
}
#[tokio::test]
async fn builder_overrides_config_for_model_endpoint_key_mode_and_reasoning() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/chat/completions"))
.and(body_string_contains("\"model\":\"test-model\""))
.and(body_string_contains("\"reasoning_effort\":\"low\""))
.respond_with(content_sse("Overrides worked."))
.expect(1)
.mount(&server)
.await;
let root = tempdir().unwrap();
let cwd = tempdir().unwrap();
write_test_config(root.path(), "https://wrong.example/v1");
let env_name = "CASSADY_EMBEDDING_TEST_KEY";
let old = std::env::var(env_name).ok();
std::env::set_var(env_name, "test-key-from-env");
let session = SessionBuilder::new()
.config_root(root.path())
.cwd(cwd.path())
.access_mode(AccessMode::WorkspaceEdit)
.model("test-model")
.base_url(server.uri())
.api_key_env(env_name)
.reasoning_effort(ReasoningEffort::Low)
.build()
.await
.unwrap();
assert_eq!(session.access_mode(), AccessMode::WorkspaceEdit);
assert_eq!(session.reasoning_effort(), ReasoningEffort::Low);
let mut turn = session.start_turn("check overrides").await.unwrap();
while let Some(event) = turn.next_event().await.unwrap() {
if matches!(event, Event::Finished) {
break;
}
}
let _session = turn.finish().await.unwrap();
if let Some(old) = old {
std::env::set_var(env_name, old);
} else {
std::env::remove_var(env_name);
}
}
#[tokio::test]
async fn embedded_session_can_resume_existing_conversation() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/chat/completions"))
.respond_with(content_sse("First turn."))
.expect(1)
.mount(&server)
.await;
let root = tempdir().unwrap();
let cwd = tempdir().unwrap();
write_test_config(root.path(), &server.uri());
let session = SessionBuilder::new()
.config_root(root.path())
.cwd(cwd.path())
.access_mode(AccessMode::ReadOnly)
.build()
.await
.unwrap();
let mut turn = session.start_turn("first").await.unwrap();
while let Some(event) = turn.next_event().await.unwrap() {
if matches!(event, Event::Finished) {
break;
}
}
let session = turn.finish().await.unwrap();
let id = session.id().to_string();
let record_count = session.records().len();
let resumed = SessionBuilder::new()
.config_root(root.path())
.cwd(cwd.path())
.resume(&id)
.await
.unwrap();
assert_eq!(resumed.id(), id);
assert_eq!(resumed.records().len(), record_count);
assert!(resumed.resume_warning().is_none());
}
#[tokio::test]
async fn embedded_approval_flow_can_approve_shell() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/chat/completions"))
.and(body_string_contains("exit code: 0"))
.respond_with(content_sse("Approved shell."))
.with_priority(1)
.expect(1)
.mount(&server)
.await;
Mock::given(method("POST"))
.and(path("/chat/completions"))
.respond_with(tool_call_sse(
"call_shell",
"shell",
r#"{"command":"touch marker"}"#,
))
.with_priority(10)
.expect(1)
.mount(&server)
.await;
let root = tempdir().unwrap();
let cwd = tempdir().unwrap();
write_test_config(root.path(), &server.uri());
let marker = cwd.path().join("marker");
let session = SessionBuilder::new()
.config_root(root.path())
.cwd(cwd.path())
.access_mode(AccessMode::WorkspaceEdit)
.build()
.await
.unwrap();
let mut turn = session.start_turn("run shell").await.unwrap();
let mut saw_request = false;
let mut saw_resolved = false;
while let Some(event) = turn.next_event().await.unwrap() {
match event {
Event::ApprovalRequested(request) => {
saw_request = true;
assert_eq!(request.name, "shell");
assert!(!marker.exists());
turn.approve(&request.request_id).unwrap();
}
Event::ApprovalResolved { approved, .. } => {
saw_resolved = approved;
}
Event::Finished => break,
_ => {}
}
}
let session = turn.finish().await.unwrap();
assert!(saw_request);
assert!(saw_resolved);
assert!(marker.exists());
assert!(session.records().iter().any(|record| matches!(
record,
Record::Tool { name, ok, content, .. }
if name == "shell" && *ok && content.contains("exit code: 0")
)));
}
#[tokio::test]
async fn read_only_embedding_does_not_advertise_mutating_tools() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/chat/completions"))
.respond_with(content_sse("Readonly."))
.expect(1)
.mount(&server)
.await;
let root = tempdir().unwrap();
let cwd = tempdir().unwrap();
write_test_config(root.path(), &server.uri());
let session = SessionBuilder::new()
.config_root(root.path())
.cwd(cwd.path())
.access_mode(AccessMode::ReadOnly)
.build()
.await
.unwrap();
let mut turn = session.start_turn("inspect only").await.unwrap();
while let Some(event) = turn.next_event().await.unwrap() {
if matches!(event, Event::Finished) {
break;
}
}
let _session = turn.finish().await.unwrap();
let requests = server.received_requests().await.unwrap();
let body = String::from_utf8_lossy(&requests[0].body);
assert!(body.contains("\"name\":\"ls\""));
assert!(body.contains("\"name\":\"read\""));
assert!(body.contains("\"name\":\"grep\""));
assert!(!body.contains("\"name\":\"write\""));
assert!(!body.contains("\"name\":\"edit\""));
assert!(!body.contains("\"name\":\"shell\""));
}