Skip to main content

starweaver_runtime/
direct.rs

1//! Direct model and tool execution helpers.
2
3use serde::{Deserialize, Serialize};
4use starweaver_core::{ConversationId, RunId, TraceContext};
5use starweaver_model::{
6    ModelAdapter, ModelError, ModelMessage, ModelRequestContext, ModelRequestParameters,
7    ModelResponse, ModelResponseStreamEvent, ModelSettings, ToolCallPart, ToolReturnPart,
8};
9use starweaver_tools::{ToolContext, ToolRegistry};
10
11/// Options for a direct model request.
12#[derive(Clone, Debug, Default, Deserialize, PartialEq, Serialize)]
13pub struct DirectModelRequest {
14    /// Canonical request history.
15    #[serde(default, skip_serializing_if = "Vec::is_empty")]
16    pub messages: Vec<ModelMessage>,
17    /// Per-call model settings.
18    #[serde(default, skip_serializing_if = "Option::is_none")]
19    pub settings: Option<ModelSettings>,
20    /// Provider-neutral request parameters.
21    #[serde(default)]
22    pub params: ModelRequestParameters,
23    /// Run id for tracing and provider metadata.
24    #[serde(default, skip_serializing_if = "Option::is_none")]
25    pub run_id: Option<RunId>,
26    /// Conversation id for tracing and provider metadata.
27    #[serde(default, skip_serializing_if = "Option::is_none")]
28    pub conversation_id: Option<ConversationId>,
29    /// Trace correlation context.
30    #[serde(default, skip_serializing_if = "TraceContext::is_empty")]
31    pub trace_context: TraceContext,
32}
33
34impl DirectModelRequest {
35    /// Build a direct request from canonical messages.
36    #[must_use]
37    pub fn new(messages: Vec<ModelMessage>) -> Self {
38        Self {
39            messages,
40            ..Self::default()
41        }
42    }
43
44    /// Attach model settings.
45    #[must_use]
46    pub fn with_settings(mut self, settings: ModelSettings) -> Self {
47        self.settings = Some(settings);
48        self
49    }
50
51    /// Attach request parameters.
52    #[must_use]
53    pub fn with_params(mut self, params: ModelRequestParameters) -> Self {
54        self.params = params;
55        self
56    }
57
58    /// Attach run and conversation identifiers.
59    #[must_use]
60    pub fn with_ids(mut self, run_id: RunId, conversation_id: ConversationId) -> Self {
61        self.run_id = Some(run_id);
62        self.conversation_id = Some(conversation_id);
63        self
64    }
65
66    /// Attach trace correlation context.
67    #[must_use]
68    pub fn with_trace_context(mut self, trace_context: TraceContext) -> Self {
69        self.trace_context = trace_context;
70        self
71    }
72
73    fn context(&self) -> ModelRequestContext {
74        ModelRequestContext::new(
75            self.run_id.clone().unwrap_or_default(),
76            self.conversation_id.clone().unwrap_or_default(),
77        )
78        .with_trace_context(self.trace_context.clone())
79    }
80}
81
82/// Execute one model request directly through a model adapter.
83///
84/// # Errors
85///
86/// Returns an error when the model adapter fails.
87pub async fn model_request(
88    model: &dyn ModelAdapter,
89    request: DirectModelRequest,
90) -> Result<ModelResponse, ModelError> {
91    let context = request.context();
92    model
93        .request_stream_final(request.messages, request.settings, request.params, context)
94        .await
95}
96
97/// Execute one model request directly and collect canonical stream events.
98///
99/// # Errors
100///
101/// Returns an error when the model adapter fails.
102pub async fn model_request_stream(
103    model: &dyn ModelAdapter,
104    request: DirectModelRequest,
105) -> Result<Vec<ModelResponseStreamEvent>, ModelError> {
106    let context = request.context();
107    model
108        .request_stream(request.messages, request.settings, request.params, context)
109        .await
110}
111
112/// Execute one tool call directly through a tool registry.
113pub async fn tool_call(
114    tools: &ToolRegistry,
115    context: ToolContext,
116    call: &ToolCallPart,
117) -> ToolReturnPart {
118    tools.execute_call(context, call).await
119}