use serde::{Deserialize, Serialize};
use starweaver_core::{ConversationId, RunId, TraceContext};
use starweaver_model::{
ModelAdapter, ModelError, ModelMessage, ModelRequestContext, ModelRequestParameters,
ModelResponse, ModelResponseStreamEvent, ModelSettings, ToolCallPart, ToolReturnPart,
};
use starweaver_tools::{ToolContext, ToolRegistry};
#[derive(Clone, Debug, Default, Deserialize, PartialEq, Serialize)]
pub struct DirectModelRequest {
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub messages: Vec<ModelMessage>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub settings: Option<ModelSettings>,
#[serde(default)]
pub params: ModelRequestParameters,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub run_id: Option<RunId>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub conversation_id: Option<ConversationId>,
#[serde(default, skip_serializing_if = "TraceContext::is_empty")]
pub trace_context: TraceContext,
}
impl DirectModelRequest {
#[must_use]
pub fn new(messages: Vec<ModelMessage>) -> Self {
Self {
messages,
..Self::default()
}
}
#[must_use]
pub fn with_settings(mut self, settings: ModelSettings) -> Self {
self.settings = Some(settings);
self
}
#[must_use]
pub fn with_params(mut self, params: ModelRequestParameters) -> Self {
self.params = params;
self
}
#[must_use]
pub fn with_ids(mut self, run_id: RunId, conversation_id: ConversationId) -> Self {
self.run_id = Some(run_id);
self.conversation_id = Some(conversation_id);
self
}
#[must_use]
pub fn with_trace_context(mut self, trace_context: TraceContext) -> Self {
self.trace_context = trace_context;
self
}
fn context(&self) -> ModelRequestContext {
ModelRequestContext::new(
self.run_id.clone().unwrap_or_default(),
self.conversation_id.clone().unwrap_or_default(),
)
.with_trace_context(self.trace_context.clone())
}
}
pub async fn model_request(
model: &dyn ModelAdapter,
request: DirectModelRequest,
) -> Result<ModelResponse, ModelError> {
let context = request.context();
model
.request_stream_final(request.messages, request.settings, request.params, context)
.await
}
pub async fn model_request_stream(
model: &dyn ModelAdapter,
request: DirectModelRequest,
) -> Result<Vec<ModelResponseStreamEvent>, ModelError> {
let context = request.context();
model
.request_stream(request.messages, request.settings, request.params, context)
.await
}
pub async fn tool_call(
tools: &ToolRegistry,
context: ToolContext,
call: &ToolCallPart,
) -> ToolReturnPart {
tools.execute_call(context, call).await
}