starweaver-session 0.2.1

Durable session contracts for Starweaver
Documentation
#![allow(clippy::unwrap_used)]

use std::sync::Arc;

use serde_json::json;
use starweaver_context::AgentContext;
use starweaver_core::{ConversationId, Metadata, RunId, TraceContext};
use starweaver_runtime::{
    AgentCheckpoint, AgentExecutionNode, AgentRunState, AgentStreamEvent, AgentStreamRecord,
};

use super::*;

#[tokio::test]
async fn input_parts_are_stable_json_contracts() {
    let input = vec![
        InputPart::text("hello"),
        InputPart::url("https://example.com"),
        InputPart::command("plan", vec!["--fast".to_string()]),
    ];

    let value = serde_json::to_value(&input).unwrap();
    assert_eq!(value[0]["kind"], "text");
    assert_eq!(value[1]["kind"], "url");
    assert_eq!(value[2]["kind"], "command");
    assert_eq!(
        serde_json::from_value::<Vec<InputPart>>(value).unwrap(),
        input
    );
}

#[test]
fn deferred_tool_facades_round_trip_records_and_decisions() {
    let session_id = SessionId::from_string("session-deferred");
    let run_id = RunId::from_string("run-deferred");
    let mut record =
        DeferredToolRecord::new("deferred-1", session_id, run_id, "call-1", "slow_tool");
    record.request = json!({"query":"rust"});
    let requests = DeferredToolRequests::from_records(&[record.clone()]);
    assert_eq!(requests.requests[0].arguments["query"], "rust");
    let rebuilt = requests.requests[0].clone().into_record();
    assert_eq!(rebuilt.deferred_id, "deferred-1");
    assert_eq!(rebuilt.request["query"], "rust");

    let mut result_metadata = Metadata::default();
    result_metadata.insert("source".to_string(), json!("worker"));
    let mut result = DeferredToolResult::completed("deferred-1", json!({"answer":"ok"}));
    result.metadata = result_metadata;
    let results = DeferredToolResults::new([result.clone()]);
    assert_eq!(results.results.len(), 1);
    result.apply_to_record(&mut record);
    assert_eq!(record.status, ExecutionStatus::Completed);
    assert_eq!(record.response["answer"], "ok");
    assert_eq!(record.metadata["source"], "worker");

    let decision = ToolApprovalDecision::approved()
        .with_override_arguments(json!({"path":"safe.txt"}))
        .into_approval_decision();
    assert_eq!(decision.status, ApprovalStatus::Approved);
    assert_eq!(decision.metadata["override_arguments"]["path"], "safe.txt");

    let denied = ToolApprovalDecision::denied("unsafe").into_approval_decision();
    assert_eq!(denied.status, ApprovalStatus::Denied);
    assert_eq!(denied.reason.as_deref(), Some("unsafe"));
}

#[test]
fn hitl_records_are_derived_from_tool_return_metadata() {
    let session_id = SessionId::from_string("session-hitl");
    let run_id = RunId::from_string("run-hitl");
    let trace_context = TraceContext::from_trace_id("trace-hitl");
    let mut approval_metadata = Metadata::default();
    approval_metadata.insert("control_flow".to_string(), json!("approval_required"));
    approval_metadata.insert("approval".to_string(), json!({"command":"rm -rf target"}));

    let approval_input = ToolReturnRecordInput::new(
        &session_id,
        &run_id,
        "call-approval",
        "shell",
        &approval_metadata,
    )
    .with_trace_context(&trace_context)
    .with_policy(json!("defer"));
    let approval = ApprovalRecord::from_tool_return(&approval_input).unwrap();
    assert_eq!(approval.approval_id, "approval_run-hitl_call-approval");
    assert_eq!(approval.action_name, "shell");
    assert_eq!(approval.request["command"], "rm -rf target");
    assert_eq!(approval.status, ApprovalStatus::Pending);
    assert_eq!(approval.metadata["policy"], "defer");
    assert_eq!(
        approval.trace_context.trace_id.as_deref(),
        Some("trace-hitl")
    );

    let mut deferred_metadata = Metadata::default();
    deferred_metadata.insert("control_flow".to_string(), json!("call_deferred"));
    deferred_metadata.insert("deferred".to_string(), json!({"url":"https://example.com"}));
    let deferred_input = ToolReturnRecordInput::new(
        &session_id,
        &run_id,
        "call-deferred",
        "fetch",
        &deferred_metadata,
    )
    .with_policy(json!("prompt"));
    let deferred = DeferredToolRecord::from_tool_return(&deferred_input).unwrap();
    assert_eq!(deferred.deferred_id, "deferred_run-hitl_call-deferred");
    assert_eq!(deferred.tool_name, "fetch");
    assert_eq!(deferred.request["url"], "https://example.com");
    assert_eq!(deferred.status, ExecutionStatus::Waiting);
    assert_eq!(deferred.metadata["policy"], "prompt");

    let ignored_metadata = Metadata::default();
    let ignored = ToolReturnRecordInput::new(
        &session_id,
        &run_id,
        "call-normal",
        "read",
        &ignored_metadata,
    );
    assert!(ApprovalRecord::from_tool_return(&ignored).is_none());
    assert!(DeferredToolRecord::from_tool_return(&ignored).is_none());
}

#[tokio::test]
async fn in_memory_store_saves_session_runs_and_resume_snapshot() {
    let store = InMemorySessionStore::new();
    let session_id = SessionId::from_string("session-1");
    let run_id = RunId::from_string("run-1");
    let conversation_id = ConversationId::from_string("conv-1");
    let mut session = SessionRecord::new(session_id.clone());
    session.profile = Some("default".to_string());
    session.workspace = Some("workspace".to_string());
    session.state = AgentContext::default().export_state();
    session.trace_context = TraceContext::from_trace_id("trace-1");
    store.save_session(session).await.unwrap();

    let mut run = RunRecord::new(session_id.clone(), run_id.clone(), conversation_id.clone());
    run.input = vec![InputPart::text("hello")];
    run.trace_context = TraceContext::from_trace_id("trace-run");
    store.append_run(run).await.unwrap();
    store
        .update_run_status(&session_id, &run_id, RunStatus::Running, None)
        .await
        .unwrap();

    let mut run_state = AgentRunState::new(run_id.clone(), conversation_id);
    run_state.run_step = 1;
    let checkpoint =
        AgentCheckpoint::new(AgentExecutionNode::ModelResponse, &run_state).with_stream_cursor(0);
    let checkpoint_id = checkpoint.checkpoint_id.clone();
    store
        .append_checkpoint(&session_id, checkpoint)
        .await
        .unwrap();
    store
        .append_stream_records(
            &session_id,
            &run_id,
            vec![
                AgentStreamRecord::new(0, AgentStreamEvent::ModelRequest { step: 0 }),
                AgentStreamRecord::new(
                    1,
                    AgentStreamEvent::RunComplete {
                        run_id: run_id.clone(),
                        output: "ok".to_string(),
                    },
                ),
            ],
        )
        .await
        .unwrap();
    store
        .append_approval(ApprovalRecord::new(
            "approval-1",
            session_id.clone(),
            run_id.clone(),
            "call-1",
            "shell",
        ))
        .await
        .unwrap();
    store
        .append_deferred_tool(DeferredToolRecord::new(
            "deferred-1",
            session_id.clone(),
            run_id.clone(),
            "call-2",
            "search",
        ))
        .await
        .unwrap();
    store
        .save_stream_cursor(
            &session_id,
            &run_id,
            StreamCursorRef::new("display", "run:run-1", 7),
        )
        .await
        .unwrap();

    let snapshot = store.resume_snapshot(&session_id, &run_id).await.unwrap();
    let trace = store.compact_run_trace(&session_id, &run_id).await.unwrap();
    let session_trace = store.compact_session_trace(&session_id).await.unwrap();

    assert_eq!(
        snapshot.latest_checkpoint.unwrap().checkpoint_id,
        checkpoint_id.clone()
    );
    assert_eq!(snapshot.stream_records.len(), 1);
    assert_eq!(snapshot.stream_records[0].sequence, 1);
    assert_eq!(snapshot.approvals.len(), 1);
    assert_eq!(snapshot.deferred_tools.len(), 1);
    assert!(snapshot
        .stream_cursors
        .iter()
        .any(|cursor| cursor.family == "display"));
    assert_eq!(trace.checkpoints, vec![checkpoint_id]);
    assert_eq!(trace.approvals, 1);
    assert_eq!(trace.deferred_tools, 1);
    assert_eq!(trace.stream_cursor, Some(1));
    assert_eq!(trace.trace_context.trace_id.as_deref(), Some("trace-run"));
    assert_eq!(session_trace.runs, 1);
    assert_eq!(session_trace.profile.as_deref(), Some("default"));
}

#[tokio::test]
async fn list_sessions_filters_and_orders_by_update_time() {
    let store = InMemorySessionStore::new();
    let mut first = SessionRecord::new(SessionId::from_string("session-a"));
    first.profile = Some("default".to_string());
    first.workspace = Some("repo-a".to_string());
    let mut second = SessionRecord::new(SessionId::from_string("session-b"));
    second.profile = Some("research".to_string());
    second.workspace = Some("repo-a".to_string());
    store.save_session(first).await.unwrap();
    store.save_session(second).await.unwrap();

    let listed = store
        .list_sessions(SessionFilter {
            workspace: Some("repo-a".to_string()),
            limit: Some(1),
            ..SessionFilter::default()
        })
        .await
        .unwrap();
    assert_eq!(listed.len(), 1);
    assert_eq!(listed[0].session_id.as_str(), "session-b");

    let filtered = store
        .list_sessions(SessionFilter {
            profile: Some("default".to_string()),
            ..SessionFilter::default()
        })
        .await
        .unwrap();
    assert_eq!(filtered.len(), 1);
    assert_eq!(filtered[0].session_id.as_str(), "session-a");
}

#[tokio::test]
async fn append_stream_records_is_idempotent_by_sequence() {
    let store = InMemorySessionStore::new();
    let session_id = SessionId::from_string("session-stream");
    let run_id = RunId::from_string("run-stream");
    store
        .save_session(SessionRecord::new(session_id.clone()))
        .await
        .unwrap();
    store
        .append_run(RunRecord::new(
            session_id.clone(),
            run_id.clone(),
            ConversationId::from_string("conv-stream"),
        ))
        .await
        .unwrap();
    store
        .append_stream_records(
            &session_id,
            &run_id,
            vec![
                AgentStreamRecord::new(1, AgentStreamEvent::ModelRequest { step: 1 }),
                AgentStreamRecord::new(0, AgentStreamEvent::ModelRequest { step: 0 }),
                AgentStreamRecord::new(1, AgentStreamEvent::ModelRequest { step: 1 }),
            ],
        )
        .await
        .unwrap();

    let replay = store
        .replay_stream_records_after(&session_id, &run_id, Some(0))
        .await
        .unwrap();
    assert_eq!(replay.len(), 1);
    assert_eq!(replay[0].sequence, 1);
}

#[tokio::test]
async fn in_memory_store_rejects_orphan_child_records() {
    let store = InMemorySessionStore::new();
    let session_id = SessionId::from_string("missing-session");
    let run_id = RunId::from_string("missing-run");
    let run = RunRecord::new(
        session_id.clone(),
        run_id.clone(),
        ConversationId::from_string("missing-conv"),
    );
    assert!(matches!(
        store.append_run(run).await,
        Err(SessionStoreError::NotFound(_))
    ));
    assert!(matches!(
        store
            .append_stream_records(
                &session_id,
                &run_id,
                vec![AgentStreamRecord::new(
                    0,
                    AgentStreamEvent::ModelRequest { step: 0 },
                )],
            )
            .await,
        Err(SessionStoreError::NotFound(_))
    ));
}

#[tokio::test]
async fn records_round_trip_through_json() {
    let session_id = SessionId::from_string("session-json");
    let run_id = RunId::from_string("run-json");
    let mut run = RunRecord::new(session_id, run_id, ConversationId::from_string("conv-json"));
    run.input = vec![InputPart::text("hello")];
    run.structured_output = json!({"ok": true});
    let mut metadata = Metadata::default();
    metadata.insert("source".to_string(), json!("test"));
    run.metadata = metadata;

    let value = serde_json::to_value(&run).unwrap();
    assert_eq!(value["status"], "queued");
    let decoded = serde_json::from_value::<RunRecord>(value).unwrap();
    assert_eq!(decoded, run);
}

#[tokio::test]
async fn session_store_executor_persists_runtime_checkpoints() {
    let store = Arc::new(InMemorySessionStore::new());
    let session_id = SessionId::from_string("session-executor");
    let session = SessionRecord::new(session_id.clone());
    store.save_session(session).await.unwrap();
    let executor = Arc::new(SessionStoreExecutor::new(store.clone(), session_id.clone()));
    let run_id = RunId::from_string("run-executor");
    let conversation_id = ConversationId::from_string("conv-executor");
    store
        .append_run(RunRecord::new(
            session_id.clone(),
            run_id.clone(),
            conversation_id.clone(),
        ))
        .await
        .unwrap();

    let mut state = AgentRunState::new(run_id.clone(), conversation_id);
    state.run_step = 2;
    starweaver_runtime::AgentExecutor::checkpoint(
        executor.as_ref(),
        AgentCheckpoint::new(AgentExecutionNode::ToolReturn, &state),
    )
    .await
    .unwrap();

    assert_eq!(
        store
            .compact_run_trace(&session_id, &run_id)
            .await
            .unwrap()
            .checkpoints
            .len(),
        1
    );
}