#![allow(dead_code)]
#![allow(clippy::unwrap_used, clippy::expect_used, clippy::panic)]
use std::sync::{Arc, Mutex, OnceLock};
use std::time::Duration;
use async_trait::async_trait;
use chrono::{DateTime, Utc};
use turnframe_core::case::{CaseKey, CaseRef, Versioned};
use turnframe_core::command::CommandBatch;
use turnframe_core::error::{ExecutionError, StoreError};
use turnframe_core::event::{Commit, CommittedEvent};
use turnframe_core::flow::{WorkflowExecutor, WorkflowRegistry};
use turnframe_core::ids::{
AccountId, CaseId, CaseRevision, CommandId, ConversationId, EventId, InteractionId,
TargetToken, TurnId,
};
use turnframe_core::interaction::Interaction;
use turnframe_core::policy::PolicySnapshot;
use turnframe_core::turn::ActorContext;
use turnframe_core::understanding::Understanding;
use turnframe_eval::corpus::EvalItem;
use turnframe_eval::runner::{EvalHarness, HarnessError, PreparedRun, SampleIndex};
use turnframe_provider::provider::ModelProvider;
use turnframe_provider::router::ProviderPool;
use turnframe_runtime::config::OrchestratorConfig;
use turnframe_runtime::orchestrator::{CaseCandidate, FixedTurnClock, Orchestrator};
use turnframe_runtime::resolve::{AuthorizedCase, TargetResolver};
use turnframe_store::events::{EventBatch, EventJournalWriter};
use turnframe_test::providers::{ScriptedProvider, ScriptedProviderBuilder, ScriptedUnderstanding};
use turnframe_test::stores::FakeStores;
use turnframe_test::workflows::InMemoryExecutor;
use turnframe_test::workflows::claim::{ClaimState, ClaimWorkflow};
use turnframe_test::workflows::traveler::{TravelerState, TravelerWorkflow};
use turnframe_test::workflows::trip::{TripState, TripWorkflow};
pub const ACCOUNT: &str = "aurora";
pub const TRIP: &str = "trip";
pub const TRAVELER: &str = "traveler";
pub const CLAIM: &str = "claim";
#[must_use]
pub fn account() -> AccountId {
AccountId::from(ACCOUNT)
}
#[must_use]
pub fn now() -> DateTime<Utc> {
DateTime::from_timestamp(1_700_000_000, 0).expect("a valid fixed instant")
}
#[must_use]
pub fn token_for(turn_id: TurnId, workflow: &str, case_id: &str) -> TargetToken {
TargetResolver::builder(account(), turn_id)
.candidate(AuthorizedCase::new(
CaseRef::new(workflow, case_id, CaseRevision::ZERO),
"label",
))
.build()
.token_map()
.token_for(&CaseKey::new(workflow, case_id))
.cloned()
.expect("a token was issued for the case")
}
#[must_use]
pub fn turn_id_for(sample: SampleIndex) -> TurnId {
TurnId::from(uuid::Uuid::from_u128(u128::from(sample.index()) + 1))
}
#[must_use]
pub fn narrating() -> ScriptedProviderBuilder {
ScriptedProvider::builder("scripted", "model-1")
.acknowledging("Right, here is where that leaves things.")
}
#[must_use]
pub fn scripted(understanding: Understanding) -> Scripted {
Scripted::narrated_by(narrating().build_shared()).understood_as(understanding)
}
#[derive(Debug, Clone)]
pub struct Scripted {
pub understander: Arc<ScriptedUnderstanding>,
pub provider: Arc<ScriptedProvider>,
}
impl Scripted {
#[must_use]
pub fn narrated_by(provider: Arc<ScriptedProvider>) -> Self {
Self {
understander: Arc::new(ScriptedUnderstanding::new()),
provider,
}
}
#[must_use]
pub fn understood_as(self, understanding: Understanding) -> Self {
self.understander.push(understanding);
self
}
}
pub struct SharedExecutor<W: turnframe_test::workflows::PureWorkflow>(pub Arc<InMemoryExecutor<W>>);
impl<W: turnframe_test::workflows::PureWorkflow> Clone for SharedExecutor<W> {
fn clone(&self) -> Self {
Self(Arc::clone(&self.0))
}
}
#[async_trait]
impl<W: turnframe_test::workflows::PureWorkflow> WorkflowExecutor<W> for SharedExecutor<W> {
async fn load(
&self,
account: &AccountId,
case_id: &CaseId,
) -> Result<Versioned<Option<W::State>>, StoreError> {
self.0.load(account, case_id).await
}
async fn execute(
&self,
batch: CommandBatch<W::Command>,
) -> Result<Commit<W::State, W::Event>, ExecutionError> {
self.0.execute(batch).await
}
}
type ScriptFn = dyn Fn(&EvalItem, SampleIndex, TurnId) -> Scripted + Send + Sync;
type ProviderFn = dyn Fn(&EvalItem, SampleIndex, TurnId) -> Arc<dyn ModelProvider> + Send + Sync;
enum Script {
Scripted(Arc<ScriptFn>),
Live(Arc<ProviderFn>),
}
#[derive(Debug, Default)]
pub struct ConcurrencyGauge {
inner: Mutex<(usize, usize)>,
}
impl ConcurrencyGauge {
#[must_use]
pub fn new() -> Self {
Self::default()
}
#[must_use]
pub fn peak(&self) -> usize {
self.inner.lock().unwrap().1
}
fn enter(&self) {
let mut guard = self.inner.lock().unwrap();
guard.0 += 1;
guard.1 = guard.1.max(guard.0);
}
fn exit(&self) {
self.inner.lock().unwrap().0 -= 1;
}
}
pub struct SampleHarness {
script: Script,
config: OrchestratorConfig,
prior_events: usize,
gauge: Option<Arc<ConcurrencyGauge>>,
}
impl SampleHarness {
#[must_use]
pub fn new(
script: impl Fn(&EvalItem, SampleIndex, TurnId) -> Scripted + Send + Sync + 'static,
) -> Self {
Self::with(Script::Scripted(Arc::new(script)))
}
#[must_use]
pub fn with_provider(
provider: impl Fn(&EvalItem, SampleIndex, TurnId) -> Arc<dyn ModelProvider>
+ Send
+ Sync
+ 'static,
) -> Self {
Self::with(Script::Live(Arc::new(provider)))
}
fn with(script: Script) -> Self {
Self {
script,
config: OrchestratorConfig::conservative(),
prior_events: 0,
gauge: None,
}
}
#[must_use]
pub const fn with_prior_events(mut self, count: usize) -> Self {
self.prior_events = count;
self
}
#[must_use]
pub fn with_effort(mut self, effort: turnframe_core::effort::Effort) -> Self {
let mut levels = turnframe_runtime::effort::EffortConfig::default();
levels.default = effort;
self.config = self.config.with_effort(levels);
self
}
#[must_use]
pub fn watched_by(mut self, gauge: Arc<ConcurrencyGauge>) -> Self {
self.gauge = Some(gauge);
self
}
}
impl std::fmt::Debug for SampleHarness {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("SampleHarness").finish_non_exhaustive()
}
}
#[async_trait]
impl EvalHarness for SampleHarness {
async fn prepare(
&self,
item: &EvalItem,
sample: SampleIndex,
) -> Result<PreparedRun, HarnessError> {
let stores = FakeStores::at(now());
let trip = Arc::new(InMemoryExecutor::new(TripWorkflow::default()));
let traveler = Arc::new(InMemoryExecutor::new(TravelerWorkflow::default()));
let claim = Arc::new(InMemoryExecutor::new(ClaimWorkflow::default()));
let account = account();
let mut seeded = Vec::new();
for seed in &item.setup.cases {
let revision = CaseRevision(seed.revision);
let state = seed.state.clone();
let parsed = |error: serde_json::Error| HarnessError::setup(error.to_string());
match seed.workflow.as_str() {
TRIP => trip.seed(
&account,
&seed.case_id,
serde_json::from_value::<TripState>(state).map_err(parsed)?,
revision,
),
TRAVELER => traveler.seed(
&account,
&seed.case_id,
serde_json::from_value::<TravelerState>(state).map_err(parsed)?,
revision,
),
CLAIM => claim.seed(
&account,
&seed.case_id,
serde_json::from_value::<ClaimState>(state).map_err(parsed)?,
revision,
),
other => {
return Err(HarnessError::UnknownWorkflow {
workflow: other.to_owned(),
});
}
}
seeded.push(CaseCandidate::new(
CaseKey::new(seed.workflow.clone(), seed.case_id.clone()),
seed.label.clone(),
));
}
let registry = WorkflowRegistry::builder()
.register(TripWorkflow::default(), SharedExecutor(Arc::clone(&trip)))
.register(
TravelerWorkflow::default(),
SharedExecutor(Arc::clone(&traveler)),
)
.register(ClaimWorkflow::default(), SharedExecutor(Arc::clone(&claim)))
.build()
.map_err(|error| HarnessError::setup(error.to_string()))?;
let turn_id = turn_id_for(sample);
let (provider, understander): (Arc<dyn ModelProvider>, _) = match &self.script {
Script::Scripted(script) => {
let scripted = script(item, sample, turn_id);
(scripted.provider, Some(scripted.understander))
}
Script::Live(provider) => (provider(item, sample, turn_id), None),
};
let trace = match &self.script {
Script::Live(_) => run_trace().map_err(HarnessError::setup)?,
Script::Scripted(_) => None,
};
let provider: Arc<dyn ModelProvider> = match &trace {
Some(trace) => Arc::new(turnframe_provider::trace::TracedProvider::new(
provider,
Arc::clone(trace) as _,
)),
None => provider,
};
let pool = Arc::new(
ProviderPool::builder()
.provider(provider)
.build()
.map_err(|error| HarnessError::setup(error.to_string()))?,
);
let workflows = Arc::new(registry);
let mut builder = Orchestrator::builder()
.workflows(Arc::clone(&workflows))
.providers(pool)
.stores(stores.stores().clone())
.case_directory(Arc::new(Records {
seeded,
trips: Arc::clone(&trip),
travelers: Arc::clone(&traveler),
}))
.policy(PolicySnapshot::conservative())
.clock(Arc::new(FixedTurnClock(now())))
.config(self.config.clone());
if let Some(understander) = understander {
builder = builder.understander(understander);
}
if let Some(trace) = trace {
builder = builder.trace(trace);
}
let orchestrator = builder
.build()
.map_err(|error| HarnessError::setup(error.to_string()))?;
let conversation_id = ConversationId::nil();
stores
.stores()
.conversations()
.create_conversation(turnframe_store::conversation::ConversationRecord::new(
conversation_id,
account.clone(),
now(),
))
.await
.map_err(|error| HarnessError::setup(error.to_string()))?;
seed_history(&stores, &account, conversation_id, item).await?;
for seed in &item.setup.cases {
seed_blocking_card(&stores, &workflows, &account, conversation_id, seed).await?;
seed_prior_events(&stores, &account, seed, self.prior_events).await?;
}
if let Some(gauge) = &self.gauge {
gauge.enter();
tokio::time::sleep(GAUGE_WINDOW).await;
gauge.exit();
}
let orchestrator = Arc::new(orchestrator);
Ok(PreparedRun {
workflows: Arc::clone(orchestrator.workflows()),
orchestrator,
actor: ActorContext::new(account, "u1"),
conversation_id,
turn_id,
})
}
}
const GAUGE_WINDOW: Duration = Duration::from_millis(25);
async fn seed_prior_events(
stores: &FakeStores,
account: &AccountId,
seed: &turnframe_eval::corpus::CaseSeed,
count: usize,
) -> Result<(), HarnessError> {
if count == 0 {
return Ok(());
}
let case = CaseKey::new(seed.workflow.clone(), seed.case_id.clone());
for chunk in 0..count.div_ceil(BATCH) {
let events: Vec<CommittedEvent<serde_json::Value>> = (0..BATCH)
.map(|offset| chunk * BATCH + offset)
.take_while(|index| *index < count)
.map(|index| CommittedEvent {
event_id: EventId::new(),
event_type: "trip.note_added".to_owned(),
occurred_at: now(),
payload: serde_json::json!({"index": index}),
})
.collect();
if events.is_empty() {
break;
}
EventJournalWriter::append(
stores.stores().events().as_ref(),
EventBatch::new(
account.clone(),
case.clone(),
CommandId::new(),
CaseRevision(seed.revision),
events,
),
)
.await
.map_err(|error| HarnessError::setup(error.to_string()))?;
}
Ok(())
}
const BATCH: usize = 64;
async fn seed_blocking_card(
stores: &FakeStores,
workflows: &WorkflowRegistry,
account: &AccountId,
conversation_id: ConversationId,
seed: &turnframe_eval::corpus::CaseSeed,
) -> Result<(), HarnessError> {
let registered =
workflows
.get(&seed.workflow)
.ok_or_else(|| HarnessError::UnknownWorkflow {
workflow: seed.workflow.as_str().to_owned(),
})?;
let case_ref = CaseRef::new(
seed.workflow.clone(),
seed.case_id.clone(),
CaseRevision(seed.revision),
);
let view = registered
.definition
.project(case_ref.clone(), Some(&seed.state))
.map_err(|error| HarnessError::setup(error.to_string()))?;
let Some(requirement) = view.blocking_interaction.as_ref() else {
return Ok(());
};
let spec = registered
.definition
.build_interaction(case_ref, Some(&seed.state), requirement)
.map_err(|error| HarnessError::setup(error.to_string()))?;
let card = Interaction::from_spec(
spec,
InteractionId::new(),
account.clone(),
conversation_id,
TurnId::nil(),
now(),
)
.map_err(|error| HarnessError::setup(error.to_string()))?;
stores
.stores()
.interactions()
.insert(card)
.await
.map_err(|error| HarnessError::setup(error.to_string()))
}
async fn seed_history(
stores: &FakeStores,
account: &AccountId,
conversation_id: ConversationId,
item: &EvalItem,
) -> Result<(), HarnessError> {
let conversations = stores.stores().conversations();
let count = item.setup.history.len();
for (index, exchange) in item.setup.history.iter().enumerate() {
let turn_id = TurnId::from(uuid::Uuid::from_u128(0xFFFF_0000 + index as u128));
let minutes_ago = i64::try_from(count - index).unwrap_or(i64::MAX);
let received_at = now() - chrono::TimeDelta::minutes(minutes_ago);
let input = turnframe_core::turn::TurnInput {
turn_id,
conversation_id,
actor: ActorContext::new(account.clone(), "u1"),
text: Some(exchange.user.clone()),
interaction_response: None,
attachments: Vec::new(),
origin: None,
locale: item
.turn
.locale
.clone()
.unwrap_or_else(|| turnframe_core::locale::Locale::from("en-GB")),
effort: None,
};
conversations
.append_user_turn(turnframe_store::conversation::StoredUserTurn::new(
input,
received_at,
))
.await
.map_err(|error| HarnessError::setup(error.to_string()))?;
let blocks = exchange
.assistant
.iter()
.map(|text| {
turnframe_core::response::ResponseBlock::Transition(
turnframe_core::response::GeneratedTransition {
block_id: turnframe_core::ids::BlockId::from(format!("history:{index}")),
text: text.clone(),
facts_used: Vec::new(),
},
)
})
.collect();
conversations
.append_assistant_turn(
account,
turnframe_core::response::AssistantTurn {
turn_id,
conversation_id,
blocks,
replay_token: turnframe_core::response::ReplayToken::from("history"),
subjects: Vec::new(),
expectations: Vec::new(),
done: Vec::new(),
},
)
.await
.map_err(|error| HarnessError::setup(error.to_string()))?;
conversations
.set_turn_phase(
account,
&turn_id,
turnframe_core::replay::TurnPhase::Delivered,
)
.await
.map_err(|error| HarnessError::setup(error.to_string()))?;
}
Ok(())
}
pub fn run_trace() -> Result<Option<Arc<turnframe_runtime::trace::JsonlTrace>>, String> {
static TRACE: OnceLock<Result<Option<Arc<turnframe_runtime::trace::JsonlTrace>>, String>> =
OnceLock::new();
TRACE
.get_or_init(|| {
turnframe_runtime::trace::JsonlTrace::from_environment(concat!(
env!("CARGO_MANIFEST_DIR"),
"/../../traces"
))
.map(|trace| trace.map(Arc::new))
.map_err(|error| error.to_string())
})
.clone()
}
struct Records {
seeded: Vec<CaseCandidate>,
trips: Arc<InMemoryExecutor<TripWorkflow>>,
travelers: Arc<InMemoryExecutor<TravelerWorkflow>>,
}
impl Records {
fn labelled(&self, account: &AccountId) -> Vec<CaseCandidate> {
let mut found = self.seeded.clone();
let known =
|found: &[CaseCandidate], key: &CaseKey| found.iter().any(|seed| seed.key == *key);
for (index, case_id) in self.trips.case_ids(account).into_iter().enumerate() {
let key = CaseKey::new(TRIP, case_id);
if !known(&found, &key) {
found.push(CaseCandidate::new(key, format!("Trip {}", index + 1)));
}
}
for (index, case_id) in self.travelers.case_ids(account).into_iter().enumerate() {
let name = self
.travelers
.state_of(account, &case_id)
.and_then(|state| state.full_name);
let key = CaseKey::new(TRAVELER, case_id);
if !known(&found, &key) {
let label = name.unwrap_or_else(|| format!("New traveler {}", index + 1));
found.push(CaseCandidate::new(key, label));
}
}
found
}
}
#[async_trait]
impl turnframe_runtime::orchestrator::CaseDirectory for Records {
async fn candidates(
&self,
actor: &ActorContext,
_conversation: &ConversationId,
) -> Result<Vec<CaseCandidate>, StoreError> {
Ok(self.labelled(&actor.account_id))
}
async fn find(
&self,
actor: &ActorContext,
_conversation: &ConversationId,
workflow: &turnframe_core::ids::WorkflowKey,
named: Option<&str>,
) -> Result<Vec<CaseCandidate>, StoreError> {
let Some(named) = named.map(str::to_lowercase) else {
return Ok(Vec::new());
};
Ok(self
.labelled(&actor.account_id)
.into_iter()
.filter(|candidate| {
candidate.key.workflow == *workflow
&& candidate.label.to_lowercase().contains(&named)
})
.collect())
}
}