use std::collections::{HashMap, HashSet};
use std::sync::Arc;
use async_trait::async_trait;
use uuid::Uuid;
use everruns_core::capabilities::CapabilityRegistry;
use everruns_core::command_host::{
CommandHost, CommandTurnContext, SessionCompletion, SessionCompletionError,
SessionCompletionRequest, SessionCompletionStream,
};
use everruns_core::execution_loading::{AgentStore, HarnessStore, SessionStore};
use everruns_core::file_services::{FileResolver, ResolvedFile};
use everruns_core::image_services::{ImageResolver, ResolvedImage};
use everruns_core::message::{Controls, Message, MessageRole, patch_dangling_tool_calls};
use everruns_core::message_retriever::MessageRetriever;
use everruns_core::provider_resolution::ProviderStore;
use everruns_core::runtime_context::{AssembledTurnContext, ResolvedModelExecution};
use everruns_core::session_files::SessionFileSystem;
use everruns_provider::driver_registry::{
ChatDriver, DriverRegistry, LlmCallConfig, LlmMessage, LlmMessageRole, ToolSearchConfig,
};
use everruns_provider::error::{AgentLoopError, Result};
use everruns_provider::runtime_provider::ProviderEndpoint;
use everruns_provider::typed_id::SessionId;
use everruns_provider::user_facing_error::UserFacingErrorContext;
use crate::runtime_context::{inspect_turn_context_for_session, resolve_model_execution};
pub struct StoreCommandHost {
session_id: SessionId,
harness_store: Arc<dyn HarnessStore>,
agent_store: Arc<dyn AgentStore>,
session_store: Arc<dyn SessionStore>,
message_retriever: Arc<dyn MessageRetriever>,
provider_store: Arc<dyn ProviderStore>,
capability_registry: CapabilityRegistry,
driver_registry: DriverRegistry,
image_resolver: Option<Arc<dyn ImageResolver>>,
file_resolver: Option<Arc<dyn FileResolver>>,
file_store: Option<Arc<dyn SessionFileSystem>>,
assembled: tokio::sync::OnceCell<AssembledTurnContext>,
}
impl StoreCommandHost {
#[allow(clippy::too_many_arguments)]
pub fn new(
session_id: SessionId,
harness_store: Arc<dyn HarnessStore>,
agent_store: Arc<dyn AgentStore>,
session_store: Arc<dyn SessionStore>,
message_retriever: Arc<dyn MessageRetriever>,
provider_store: Arc<dyn ProviderStore>,
capability_registry: CapabilityRegistry,
driver_registry: DriverRegistry,
) -> Self {
Self {
session_id,
harness_store,
agent_store,
session_store,
message_retriever,
provider_store,
capability_registry,
driver_registry,
image_resolver: None,
file_resolver: None,
file_store: None,
assembled: tokio::sync::OnceCell::new(),
}
}
pub fn with_image_resolver(mut self, image_resolver: Arc<dyn ImageResolver>) -> Self {
self.image_resolver = Some(image_resolver);
self
}
pub fn with_file_resolver(mut self, file_resolver: Arc<dyn FileResolver>) -> Self {
self.file_resolver = Some(file_resolver);
self
}
pub fn with_file_store(mut self, file_store: Arc<dyn SessionFileSystem>) -> Self {
self.file_store = Some(file_store);
self
}
pub fn with_assembled_context(mut self, assembled: AssembledTurnContext) -> Self {
self.assembled = tokio::sync::OnceCell::new_with(Some(assembled));
self
}
async fn assembled(&self) -> Result<&AssembledTurnContext> {
self.assembled
.get_or_try_init(|| async {
inspect_turn_context_for_session(
self.harness_store.as_ref(),
self.agent_store.as_ref(),
self.session_store.as_ref(),
self.message_retriever.as_ref(),
self.provider_store.as_ref(),
&self.capability_registry,
&self.driver_registry,
self.session_id,
&[],
self.file_store.clone(),
)
.await
})
.await
}
async fn resolve_images(&self, messages: &[Message]) -> HashMap<Uuid, ResolvedImage> {
let Some(resolver) = &self.image_resolver else {
return HashMap::new();
};
let image_ids: HashSet<Uuid> = messages
.iter()
.flat_map(everruns_core::llm_conversions::extract_image_file_ids)
.collect();
let mut resolved = HashMap::new();
for image_id in image_ids {
if let Ok(Some(image)) = resolver.resolve_image(image_id).await {
resolved.insert(image_id, image);
}
}
resolved
}
async fn resolve_files(&self, messages: &[Message]) -> HashMap<Uuid, ResolvedFile> {
let Some(resolver) = &self.file_resolver else {
return HashMap::new();
};
let file_ids: HashSet<Uuid> = messages
.iter()
.flat_map(everruns_core::llm_conversions::extract_file_ids)
.collect();
match resolver
.resolve_files(&file_ids.into_iter().collect::<Vec<_>>())
.await
{
Ok(map) => map,
Err(e) => {
tracing::warn!("Failed to resolve file attachments: {e}");
HashMap::new()
}
}
}
async fn resolve_completion_model(
&self,
controls: Option<&Controls>,
assembled: &AssembledTurnContext,
) -> std::result::Result<ResolvedModelExecution, SessionCompletionError> {
let requested = controls.and_then(|controls| controls.model_id);
if requested.is_none() || requested == assembled.resolved_model_id {
return Ok(assembled.model.clone());
}
let model_id = requested.expect("checked above");
let spec = self
.provider_store
.get_model_spec(model_id)
.await
.map_err(SessionCompletionError::InvalidRequest)?
.ok_or_else(|| {
SessionCompletionError::InvalidRequest(AgentLoopError::config(format!(
"Model not found: {model_id}"
)))
})?;
let error_context = UserFacingErrorContext::default()
.with_provider(spec.provider.to_string())
.with_model_id(spec.model.clone());
resolve_model_execution(self.provider_store.as_ref(), &self.driver_registry, spec)
.await
.map_err(|error| SessionCompletionError::Completion {
error: error.to_string(),
context: error_context,
})
}
async fn prepare_completion(
&self,
request: SessionCompletionRequest,
) -> std::result::Result<PreparedCompletion, SessionCompletionError> {
let assembled = self
.assembled()
.await
.map_err(SessionCompletionError::InvalidRequest)?;
let model = self
.resolve_completion_model(request.controls.as_ref(), assembled)
.await?;
let context = UserFacingErrorContext::default()
.with_provider(model.provider_type.to_string())
.with_model_id(model.model.clone());
let messages = patch_dangling_tool_calls(&request.messages);
let resolved_images = self.resolve_images(&messages).await;
let resolved_files = self.resolve_files(&messages).await;
let mut llm_messages: Vec<LlmMessage> = request
.system_prompts
.iter()
.filter(|prompt| !prompt.is_empty())
.map(|prompt| LlmMessage::text(LlmMessageRole::System, prompt.clone()))
.collect();
for message in &messages {
let mut llm_message =
everruns_core::llm_conversions::llm_message_from_message_with_attachments(
message,
&resolved_images,
&resolved_files,
);
if message.role == MessageRole::User
&& let Some(actor) = &message.external_actor
{
llm_message.prepend_text_prefix(&format!("[{}] ", actor.display_label()));
}
llm_messages.push(llm_message);
}
let mut builder = everruns_core::llm_conversions::llm_call_config_builder_from_agent(
&assembled.runtime_agent,
)
.model(&model.model)
.tools(vec![])
.tool_search(ToolSearchConfig {
enabled: false,
threshold: usize::MAX,
})
.previous_response_id(None)
.with_metadata("session_id", self.session_id.to_string());
if let Some(effort) = request
.controls
.as_ref()
.and_then(|controls| controls.reasoning.as_ref())
.and_then(|reasoning| reasoning.effort)
{
builder = builder.reasoning_effort(effort);
}
for (key, value) in &request.metadata {
builder = builder.with_metadata(key, value);
}
Ok(PreparedCompletion {
llm_messages,
llm_config: builder.build(),
driver: model.driver,
context,
})
}
}
struct PreparedCompletion {
llm_messages: Vec<LlmMessage>,
llm_config: LlmCallConfig,
driver: Arc<dyn ChatDriver>,
context: UserFacingErrorContext,
}
#[async_trait]
impl CommandHost for StoreCommandHost {
async fn turn_context(&self) -> Result<CommandTurnContext> {
let assembled = self.assembled().await?;
Ok(CommandTurnContext {
session_id: assembled.snapshot.session_id,
messages: assembled.messages.clone(),
system_prompt: assembled.runtime_agent.system_prompt.clone(),
model: assembled.model.model.clone(),
provider_type: assembled.model.provider_type.to_string(),
resolved_locale: assembled.resolved_locale.clone(),
})
}
async fn completion(
&self,
request: SessionCompletionRequest,
) -> std::result::Result<SessionCompletion, SessionCompletionError> {
let prepared = self.prepare_completion(request).await?;
let completion_error = |error: String| SessionCompletionError::Completion {
error,
context: prepared.context.clone(),
};
let response = prepared
.driver
.chat_completion(
&ProviderEndpoint::default(),
prepared.llm_messages,
&prepared.llm_config,
)
.await
.map_err(|error| completion_error(error.to_string()))?;
let text = response.text.trim().to_string();
if text.is_empty() {
return Err(completion_error(
"session completion returned an empty response".to_string(),
));
}
Ok(SessionCompletion { text })
}
async fn completion_stream(
&self,
request: SessionCompletionRequest,
) -> std::result::Result<SessionCompletionStream, SessionCompletionError> {
let prepared = self.prepare_completion(request).await?;
let events = prepared
.driver
.chat_completion_stream(
&ProviderEndpoint::default(),
prepared.llm_messages,
&prepared.llm_config,
)
.await
.map_err(|error| SessionCompletionError::Completion {
error: error.to_string(),
context: prepared.context.clone(),
})?;
Ok(SessionCompletionStream {
events,
context: prepared.context,
})
}
}