Skip to main content

lash_core/
direct.rs

1use crate::llm::transport::LlmTransportError;
2use crate::llm::types::{
3    AttachmentSource, LlmContentBlock, LlmEventSender, LlmJsonSchema, LlmMessage, LlmOutputSpec,
4    LlmRequest, LlmRequestScope, LlmResponse, LlmRole, LlmStreamEvent, LlmTerminalReason,
5    LlmToolChoice,
6};
7use crate::provider::{ModelCapability, ModelEffortValidationCategory, ProviderHandle};
8use crate::{LashSchema, SchemaContract};
9use lash_trace::{TraceContext, TraceError, TraceEvent, TraceSink};
10use std::sync::Arc;
11
12#[derive(Clone, Debug, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
13#[serde(rename_all = "snake_case")]
14pub enum DirectRole {
15    System,
16    User,
17    Assistant,
18}
19
20#[derive(Clone, Debug, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
21pub enum DirectPart {
22    Text(String),
23    Attachment(usize),
24}
25
26#[derive(Clone, Debug, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
27pub struct DirectMessage {
28    pub role: DirectRole,
29    pub parts: Vec<DirectPart>,
30}
31
32#[derive(Clone, Debug, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
33pub struct DirectJsonSchema {
34    pub name: String,
35    pub schema: SchemaContract,
36    pub strict: bool,
37}
38
39#[derive(Clone, Debug, PartialEq, Eq, Default, serde::Serialize, serde::Deserialize)]
40pub enum DirectOutputSpec {
41    #[default]
42    Text,
43    JsonObject,
44    JsonSchema(DirectJsonSchema),
45}
46
47#[derive(Clone, Debug, serde::Serialize, serde::Deserialize)]
48pub struct DirectRequest {
49    pub model: String,
50    #[serde(default)]
51    pub model_variant: crate::ReasoningSelection,
52    #[serde(default, skip_serializing_if = "ModelCapability::is_empty")]
53    pub model_capability: ModelCapability,
54    #[serde(default)]
55    pub messages: Vec<DirectMessage>,
56    #[serde(default)]
57    pub attachments: Vec<AttachmentSource>,
58    #[serde(default)]
59    pub output: DirectOutputSpec,
60    #[serde(default)]
61    pub generation: crate::GenerationOptions,
62    #[serde(default, skip)]
63    pub stream_events: Option<LlmEventSender>,
64    #[serde(default, skip_serializing_if = "Option::is_none")]
65    pub session_id: Option<String>,
66    #[serde(default, skip_serializing_if = "Option::is_none")]
67    pub caused_by: Option<crate::CausalRef>,
68    #[serde(default, skip_serializing_if = "Option::is_none")]
69    pub replay: Option<crate::RuntimeReplay>,
70}
71
72impl DirectRequest {
73    pub fn text(model: impl Into<String>, prompt: impl Into<String>) -> Self {
74        Self {
75            model: model.into(),
76            model_variant: crate::ReasoningSelection::ProviderDefault,
77            model_capability: ModelCapability::default(),
78            messages: vec![DirectMessage {
79                role: DirectRole::User,
80                parts: vec![DirectPart::Text(prompt.into())],
81            }],
82            attachments: Vec::new(),
83            output: DirectOutputSpec::Text,
84            generation: crate::GenerationOptions::default(),
85            stream_events: None,
86            session_id: None,
87            caused_by: None,
88            replay: None,
89        }
90    }
91
92    pub fn json(model: impl Into<String>, prompt: impl Into<String>) -> Self {
93        Self {
94            output: DirectOutputSpec::JsonObject,
95            ..Self::text(model, prompt)
96        }
97    }
98
99    pub fn json_schema(
100        model: impl Into<String>,
101        prompt: impl Into<String>,
102        schema: DirectJsonSchema,
103    ) -> Self {
104        Self {
105            output: DirectOutputSpec::JsonSchema(schema),
106            ..Self::text(model, prompt)
107        }
108    }
109
110    pub fn with_replay_key(mut self, key: impl Into<String>) -> Self {
111        self.replay = Some(crate::RuntimeReplay { key: key.into() });
112        self
113    }
114
115    pub fn with_caused_by(mut self, caused_by: crate::CausalRef) -> Self {
116        self.caused_by = Some(caused_by);
117        self
118    }
119}
120
121#[derive(Debug, thiserror::Error, Clone)]
122pub enum DirectLlmError {
123    #[error("invalid request: {message}")]
124    InvalidRequest {
125        category: ModelEffortValidationCategory,
126        message: String,
127    },
128    #[error("invalid response: {0}")]
129    InvalidResponse(String),
130    #[error("transport error: {0}")]
131    Transport(#[from] Box<LlmTransportError>),
132}
133
134/// Successful single-shot direct LLM result with the sealed provider-attempt
135/// history that produced it.
136#[derive(Clone, Debug)]
137pub struct DirectLlmResult {
138    pub response: LlmResponse,
139    pub llm_call: crate::LlmCallRecord,
140}
141
142impl std::ops::Deref for DirectLlmResult {
143    type Target = LlmResponse;
144
145    fn deref(&self) -> &Self::Target {
146        &self.response
147    }
148}
149
150impl DirectLlmResult {
151    pub fn into_response(self) -> LlmResponse {
152        self.response
153    }
154}
155
156pub struct DirectLlmClient {
157    provider: ProviderHandle,
158    trace_sink: Option<Arc<dyn TraceSink>>,
159    trace_context: TraceContext,
160    clock: Arc<dyn crate::Clock>,
161}
162
163impl DirectLlmClient {
164    pub fn new(provider: ProviderHandle) -> Self {
165        Self {
166            provider,
167            trace_sink: None,
168            trace_context: TraceContext::default(),
169            clock: Arc::new(crate::SystemClock),
170        }
171    }
172
173    pub fn with_trace_sink(mut self, sink: Option<Arc<dyn TraceSink>>) -> Self {
174        self.trace_sink = sink;
175        self
176    }
177
178    pub fn with_trace_context(mut self, context: TraceContext) -> Self {
179        self.trace_context = context;
180        self
181    }
182
183    pub fn with_clock(mut self, clock: Arc<dyn crate::Clock>) -> Self {
184        self.clock = clock;
185        self
186    }
187
188    pub fn provider(&self) -> &ProviderHandle {
189        &self.provider
190    }
191
192    pub fn provider_mut(&mut self) -> &mut ProviderHandle {
193        &mut self.provider
194    }
195
196    pub async fn complete(
197        &mut self,
198        mut request: DirectRequest,
199    ) -> Result<DirectLlmResult, DirectLlmError> {
200        // Validate the requested effort against the capability that travels
201        // with the request, and write the resolved (alias-normalized) effort
202        // back so the provider never sees an un-clamped value.
203        request.model_variant = request
204            .model_capability
205            .validate_selection(&request.model, self.provider.kind(), &request.model_variant)
206            .map_err(|error| DirectLlmError::InvalidRequest {
207                category: error.category,
208                message: error.message,
209            })?;
210
211        let output_for_validation = request.output.clone();
212        let model = request.model.clone();
213        let llm_request = build_llm_request(&self.provider, request, model);
214        let llm_call_id = if self.trace_sink.is_some() {
215            let id = uuid::Uuid::new_v4().to_string();
216            crate::trace::emit_trace(
217                &self.trace_sink,
218                &self.trace_context,
219                TraceContext::default().for_llm_call(id.clone()),
220                TraceEvent::LlmCallStarted {
221                    request: crate::trace::trace_llm_request(&llm_request),
222                },
223                self.clock.as_ref(),
224            );
225            Some(id)
226        } else {
227            None
228        };
229        match self.provider.complete(llm_request).await {
230            Ok(response) => {
231                if let Err(error) = validate_direct_output(&output_for_validation, &response) {
232                    if let Some(llm_call_id) = llm_call_id {
233                        crate::trace::emit_trace(
234                            &self.trace_sink,
235                            &self.trace_context,
236                            TraceContext::default().for_llm_call(llm_call_id),
237                            TraceEvent::LlmCallFailed {
238                                error: TraceError {
239                                    message: error.to_string(),
240                                    retryable: false,
241                                    terminal_reason: Some(
242                                        LlmTerminalReason::ProviderError.code().to_string(),
243                                    ),
244                                    code: Some("invalid_structured_output".to_string()),
245                                    raw: None,
246                                },
247                                stream_summary: None,
248                            },
249                            self.clock.as_ref(),
250                        );
251                    }
252                    return Err(error);
253                }
254                if let Some(llm_call_id) = llm_call_id {
255                    crate::trace::emit_trace(
256                        &self.trace_sink,
257                        &self.trace_context,
258                        TraceContext::default().for_llm_call(llm_call_id),
259                        TraceEvent::LlmCallCompleted {
260                            response: crate::trace::trace_llm_response(
261                                response.full_text.clone(),
262                                0,
263                                Some(response.terminal_reason),
264                                crate::trace::trace_output_parts(&response.parts),
265                            ),
266                            usage: Some(crate::trace::trace_usage_from_llm(&response.usage)),
267                            provider_usage: response.provider_usage.clone(),
268                            stream_summary: None,
269                        },
270                        self.clock.as_ref(),
271                    );
272                }
273                Ok(DirectLlmResult {
274                    response: response.response,
275                    llm_call: response.call_record,
276                })
277            }
278            Err(error) => {
279                if let Some(llm_call_id) = llm_call_id {
280                    crate::trace::emit_trace(
281                        &self.trace_sink,
282                        &self.trace_context,
283                        TraceContext::default().for_llm_call(llm_call_id),
284                        TraceEvent::LlmCallFailed {
285                            error: TraceError {
286                                message: error.message.clone(),
287                                retryable: error.retryable,
288                                terminal_reason: Some(error.terminal_reason.code().to_string()),
289                                code: error.code.clone(),
290                                raw: error.raw.as_deref().cloned(),
291                            },
292                            stream_summary: None,
293                        },
294                        self.clock.as_ref(),
295                    );
296                }
297                Err(DirectLlmError::from(Box::new(error.error)))
298            }
299        }
300    }
301}
302
303pub(crate) fn build_llm_request(
304    provider: &ProviderHandle,
305    request: DirectRequest,
306    model: String,
307) -> LlmRequest {
308    let stream_events = transport_stream_events_for_direct(provider, request.stream_events);
309    let DirectRequest {
310        model: _,
311        model_variant,
312        model_capability,
313        messages,
314        attachments,
315        output,
316        generation,
317        stream_events: _,
318        session_id,
319        caused_by: _,
320        replay: _,
321    } = request;
322
323    let output_spec = match output {
324        DirectOutputSpec::Text => None,
325        DirectOutputSpec::JsonObject => Some(LlmOutputSpec::JsonObject),
326        DirectOutputSpec::JsonSchema(schema) => Some(LlmOutputSpec::JsonSchema(LlmJsonSchema {
327            name: schema.name,
328            schema: schema.schema,
329            strict: schema.strict,
330        })),
331    };
332
333    let mut llm_messages = Vec::new();
334    for message in messages {
335        let role = match message.role {
336            DirectRole::System => LlmRole::System,
337            DirectRole::User => LlmRole::User,
338            DirectRole::Assistant => LlmRole::Assistant,
339        };
340        let mut blocks: Vec<LlmContentBlock> = Vec::new();
341        for part in message.parts {
342            match part {
343                DirectPart::Text(text) => {
344                    if !text.is_empty() {
345                        blocks.push(LlmContentBlock::Text {
346                            text: text.into(),
347                            response_meta: None,
348                            cache_breakpoint: false,
349                        });
350                    }
351                }
352                DirectPart::Attachment(idx) => {
353                    blocks.push(LlmContentBlock::Attachment {
354                        attachment_idx: idx,
355                    });
356                }
357            }
358        }
359        if !blocks.is_empty() {
360            llm_messages.push(LlmMessage::new(role, blocks));
361        }
362    }
363
364    let scope = match session_id {
365        Some(session_id) => LlmRequestScope::new(
366            session_id.clone(),
367            format!("{session_id}:frame:direct"),
368            format!("{session_id}:direct"),
369        ),
370        None => {
371            let request_id = uuid::Uuid::new_v4().to_string();
372            LlmRequestScope::new(
373                format!("direct:{request_id}"),
374                format!("direct:{request_id}:frame"),
375                request_id,
376            )
377        }
378    };
379
380    LlmRequest {
381        model,
382        messages: llm_messages,
383        attachments,
384        resolved_stored: Default::default(),
385        tools: Vec::new().into(),
386        tool_choice: LlmToolChoice::None,
387        model_variant,
388        model_capability,
389        generation,
390        scope,
391        output_spec,
392        stream_events,
393        provider_trace: None,
394    }
395}
396
397fn validate_direct_output(
398    output: &DirectOutputSpec,
399    response: &LlmResponse,
400) -> Result<(), DirectLlmError> {
401    let DirectOutputSpec::JsonSchema(schema) = output else {
402        return Ok(());
403    };
404    let parsed: serde_json::Value = serde_json::from_str(response.full_text.trim())
405        .map_err(|err| DirectLlmError::InvalidResponse(format!("expected JSON: {err}")))?;
406    LashSchema::new(schema.schema.canonical().clone())
407        .validate(&parsed)
408        .map_err(DirectLlmError::InvalidResponse)
409}
410
411fn transport_stream_events_for_direct(
412    provider: &ProviderHandle,
413    requested: Option<LlmEventSender>,
414) -> Option<LlmEventSender> {
415    if requested.is_some() {
416        return requested;
417    }
418    if provider.requires_streaming() {
419        Some(LlmEventSender::new(|_event: LlmStreamEvent| {}))
420    } else {
421        None
422    }
423}
424
425#[cfg(test)]
426mod tests {
427    use super::*;
428    use crate::llm::types::{LlmOutputPart, LlmTerminalReason, LlmUsage};
429    use crate::provider::{ProviderOptions, ProviderReliability};
430    use crate::testing::TestProvider;
431    use serde_json::json;
432    use std::sync::{Arc, Mutex};
433
434    #[test]
435    fn json_schema_request_preserves_output_schema() {
436        let schema = DirectJsonSchema {
437            name: "answer_shape".to_string(),
438            schema: json!({
439                "type": "object",
440                "properties": {
441                    "answer": { "type": "string" }
442                },
443                "required": ["answer"]
444            })
445            .into(),
446            strict: true,
447        };
448
449        let request = DirectRequest::json_schema("model-a", "return json", schema.clone());
450
451        assert_eq!(
452            request.output,
453            DirectOutputSpec::JsonSchema(schema),
454            "DirectRequest::json_schema must carry the requested output schema"
455        );
456    }
457
458    #[test]
459    fn direct_client_provider_accessors_expose_owned_provider_handle() {
460        let provider = TestProvider::builder()
461            .kind("direct-accessor-provider")
462            .serialize_config(|| json!({"provider": "owned"}))
463            .build()
464            .into_handle();
465        let mut client = DirectLlmClient::new(provider);
466
467        assert_eq!(client.provider().kind(), "direct-accessor-provider");
468        assert_eq!(
469            client.provider().to_spec().config,
470            json!({"provider": "owned"})
471        );
472
473        let options = ProviderOptions {
474            reliability: ProviderReliability::default().max_attempts(7),
475            max_output_tokens: Some(123),
476            ..Default::default()
477        };
478        client.provider_mut().set_options(options.clone());
479
480        assert_eq!(client.provider().options(), options);
481    }
482
483    #[tokio::test]
484    async fn direct_client_complete_delegates_to_provider_and_returns_response() {
485        let captured_request: Arc<Mutex<Option<LlmRequest>>> = Arc::new(Mutex::new(None));
486        let captured_for_provider = Arc::clone(&captured_request);
487        let provider = TestProvider::builder()
488            .kind("direct-complete-provider")
489            .complete(move |request| {
490                let captured_for_provider = Arc::clone(&captured_for_provider);
491                async move {
492                    *captured_for_provider.lock().expect("capture lock") = Some(request);
493                    Ok(LlmResponse {
494                        full_text: "provider delegated response".to_string(),
495                        parts: vec![LlmOutputPart::Text {
496                            text: "provider delegated response".to_string(),
497                            response_meta: None,
498                        }],
499                        usage: LlmUsage {
500                            input_tokens: 11,
501                            output_tokens: 3,
502                            ..Default::default()
503                        },
504                        terminal_reason: LlmTerminalReason::Stop,
505                        response_metadata: Default::default(),
506                        ..Default::default()
507                    })
508                }
509            })
510            .build()
511            .into_handle();
512        let mut client = DirectLlmClient::new(provider);
513        let mut request = DirectRequest::json("direct-model", "answer as json");
514        request.session_id = Some("direct-session".to_string());
515
516        let response = client
517            .complete(request)
518            .await
519            .expect("direct completion should delegate");
520
521        assert_eq!(response.full_text, "provider delegated response");
522        assert_eq!(response.llm_call.attempts.len(), 1);
523        let captured = captured_request
524            .lock()
525            .expect("capture lock")
526            .clone()
527            .expect("provider should receive a request");
528        assert_eq!(captured.model, "direct-model");
529        assert_eq!(captured.scope.session_id, "direct-session");
530        assert_eq!(captured.scope.agent_frame_id, "direct-session:frame:direct");
531        assert_eq!(captured.scope.request_id, "direct-session:direct");
532        assert!(matches!(
533            captured.output_spec,
534            Some(LlmOutputSpec::JsonObject)
535        ));
536        assert_eq!(captured.messages.len(), 1);
537    }
538
539    #[tokio::test]
540    async fn direct_client_validates_json_schema_output_against_canonical_schema() {
541        let provider = TestProvider::builder()
542            .kind("direct-validation-provider")
543            .complete(|_request| async {
544                Ok(LlmResponse {
545                    full_text: r#"{"items":[]}"#.to_string(),
546                    terminal_reason: LlmTerminalReason::Stop,
547                    response_metadata: Default::default(),
548                    ..Default::default()
549                })
550            })
551            .build()
552            .into_handle();
553        let mut client = DirectLlmClient::new(provider);
554        let request = DirectRequest::json_schema(
555            "direct-model",
556            "return items",
557            DirectJsonSchema {
558                name: "items_result".to_string(),
559                schema: json!({
560                    "type": "object",
561                    "required": ["items"],
562                    "properties": {
563                        "items": {
564                            "type": "array",
565                            "minItems": 1,
566                            "items": { "type": "string" }
567                        }
568                    }
569                })
570                .into(),
571                strict: true,
572            },
573        );
574
575        let err = client
576            .complete(request)
577            .await
578            .expect_err("empty items must fail canonical validation");
579
580        assert!(matches!(err, DirectLlmError::InvalidResponse(_)));
581        let error = err.to_string();
582        assert!(
583            error.contains("items") && error.contains("[] has less than 1 item"),
584            "{error}"
585        );
586    }
587
588    fn reasoning_capability() -> ModelCapability {
589        ModelCapability {
590            reasoning: Some(crate::ReasoningCapability {
591                efforts: ["low", "medium", "high", "max"]
592                    .into_iter()
593                    .map(String::from)
594                    .collect(),
595                aliases: std::collections::BTreeMap::from([(
596                    "xhigh".to_string(),
597                    "max".to_string(),
598                )]),
599                ..Default::default()
600            }),
601            cache_control: None,
602            stream_termination: None,
603        }
604    }
605
606    #[tokio::test]
607    async fn direct_client_rejects_unsupported_effort_before_provider_call() {
608        let called = Arc::new(Mutex::new(false));
609        let called_for_provider = Arc::clone(&called);
610        let provider = TestProvider::builder()
611            .kind("direct-reject")
612            .complete(move |_request| {
613                let called = Arc::clone(&called_for_provider);
614                async move {
615                    *called.lock().expect("called lock") = true;
616                    Ok(LlmResponse::default())
617                }
618            })
619            .build()
620            .into_handle();
621        let mut client = DirectLlmClient::new(provider);
622
623        let mut request = DirectRequest::text("direct-model", "hi");
624        request.model_variant = crate::ReasoningSelection::Effort("turbo".to_string());
625        request.model_capability = reasoning_capability();
626
627        let err = client
628            .complete(request)
629            .await
630            .expect_err("unsupported effort must be rejected");
631        assert!(matches!(
632            err,
633            DirectLlmError::InvalidRequest {
634                category: ModelEffortValidationCategory::UnsupportedEffort,
635                ..
636            }
637        ));
638        assert!(err.to_string().contains("Unsupported effort `turbo`"));
639        assert!(
640            !*called.lock().expect("called lock"),
641            "the provider must not be called when the effort is rejected"
642        );
643    }
644
645    #[tokio::test]
646    async fn direct_client_normalizes_alias_effort_into_outgoing_request() {
647        let captured: Arc<Mutex<Option<crate::ReasoningSelection>>> = Arc::new(Mutex::new(None));
648        let captured_for_provider = Arc::clone(&captured);
649        let provider = TestProvider::builder()
650            .kind("direct-alias")
651            .complete(move |request| {
652                let captured = Arc::clone(&captured_for_provider);
653                async move {
654                    *captured.lock().expect("capture lock") = Some(request.model_variant.clone());
655                    Ok(LlmResponse {
656                        full_text: "ok".to_string(),
657                        terminal_reason: LlmTerminalReason::Stop,
658                        response_metadata: Default::default(),
659                        ..Default::default()
660                    })
661                }
662            })
663            .build()
664            .into_handle();
665        let mut client = DirectLlmClient::new(provider);
666
667        let mut request = DirectRequest::text("direct-model", "hi");
668        request.model_variant = crate::ReasoningSelection::Effort("XHigh".to_string());
669        request.model_capability = reasoning_capability();
670
671        client.complete(request).await.expect("completion");
672        let seen = captured
673            .lock()
674            .expect("capture lock")
675            .clone()
676            .expect("provider must be called");
677        assert_eq!(
678            seen,
679            crate::ReasoningSelection::Effort("max".to_string()),
680            "alias `XHigh` must clamp to canonical `max` before the provider sees the request"
681        );
682    }
683
684    #[tokio::test]
685    async fn direct_client_rejects_effort_when_model_is_not_configurable() {
686        let provider = TestProvider::builder()
687            .kind("direct-not-configurable")
688            .complete(|_request| async { Ok(LlmResponse::default()) })
689            .build()
690            .into_handle();
691        let mut client = DirectLlmClient::new(provider);
692
693        let mut request = DirectRequest::text("direct-model", "hi");
694        request.model_variant = crate::ReasoningSelection::Effort("high".to_string());
695        // No capability: the model exposes no configurable effort.
696
697        let err = client
698            .complete(request)
699            .await
700            .expect_err("effort on a non-configurable model must be rejected");
701        assert!(matches!(
702            err,
703            DirectLlmError::InvalidRequest {
704                category: ModelEffortValidationCategory::EffortNotConfigurable,
705                ..
706            }
707        ));
708    }
709
710    #[tokio::test]
711    async fn direct_client_rejects_missing_mandatory_effort() {
712        let provider = TestProvider::builder()
713            .kind("direct-mandatory")
714            .complete(|_request| async { Ok(LlmResponse::default()) })
715            .build()
716            .into_handle();
717        let mut client = DirectLlmClient::new(provider);
718
719        let mut capability = reasoning_capability();
720        capability.reasoning.as_mut().expect("reasoning").mandatory = true;
721        let mut request = DirectRequest::text("direct-model", "hi");
722        request.model_capability = capability;
723        // No model_variant supplied, but the model requires one.
724
725        let err = client
726            .complete(request)
727            .await
728            .expect_err("missing mandatory effort must be rejected");
729        assert!(matches!(
730            err,
731            DirectLlmError::InvalidRequest {
732                category: ModelEffortValidationCategory::EffortRequired,
733                ..
734            }
735        ));
736    }
737
738    #[test]
739    fn build_llm_request_preserves_nonempty_content_and_drops_empty_messages() {
740        let provider = TestProvider::default().into_handle();
741        let request = DirectRequest {
742            model: "input-model".to_string(),
743            messages: vec![
744                DirectMessage {
745                    role: DirectRole::System,
746                    parts: vec![DirectPart::Text(String::new())],
747                },
748                DirectMessage {
749                    role: DirectRole::User,
750                    parts: vec![
751                        DirectPart::Text("hello".to_string()),
752                        DirectPart::Text(String::new()),
753                    ],
754                },
755                DirectMessage {
756                    role: DirectRole::Assistant,
757                    parts: vec![DirectPart::Attachment(2)],
758                },
759            ],
760            attachments: Vec::new(),
761            output: DirectOutputSpec::Text,
762            generation: crate::GenerationOptions::default(),
763            stream_events: None,
764            session_id: None,
765            model_variant: Default::default(),
766            model_capability: ModelCapability::default(),
767            caused_by: None,
768            replay: None,
769        };
770
771        let llm_request = build_llm_request(&provider, request, "transport-model".to_string());
772
773        assert_eq!(llm_request.model, "transport-model");
774        assert_eq!(
775            llm_request.messages.len(),
776            2,
777            "empty normalized messages must be dropped"
778        );
779        assert_eq!(llm_request.messages[0].role, LlmRole::User);
780        assert_eq!(llm_request.messages[0].blocks.len(), 1);
781        assert!(matches!(
782            &llm_request.messages[0].blocks[0],
783            LlmContentBlock::Text { text, .. } if text.as_ref() == "hello"
784        ));
785        assert_eq!(llm_request.messages[1].role, LlmRole::Assistant);
786        assert!(matches!(
787            &llm_request.messages[1].blocks[0],
788            LlmContentBlock::Attachment { attachment_idx: 2 }
789        ));
790    }
791
792    #[test]
793    fn build_llm_request_preserves_direct_stream_sender_and_adds_required_noop_sender() {
794        let captured_events: Arc<Mutex<Vec<LlmStreamEvent>>> = Arc::new(Mutex::new(Vec::new()));
795        let captured_for_sender = Arc::clone(&captured_events);
796        let requested_sender = LlmEventSender::new(move |event| {
797            captured_for_sender
798                .lock()
799                .expect("stream event lock")
800                .push(event);
801        });
802        let mut request = DirectRequest::text("model", "prompt");
803        request.stream_events = Some(requested_sender);
804        let provider = TestProvider::default().into_handle();
805
806        let llm_request = build_llm_request(&provider, request, "model".to_string());
807        let sender = llm_request
808            .stream_events
809            .expect("explicit direct stream sender must be preserved");
810        sender.send(LlmStreamEvent::Delta("delta".to_string()));
811        assert_eq!(captured_events.lock().expect("stream event lock").len(), 1);
812
813        let streaming_provider = TestProvider::builder()
814            .requires_streaming(true)
815            .build()
816            .into_handle();
817        let llm_request = build_llm_request(
818            &streaming_provider,
819            DirectRequest::text("model", "prompt"),
820            "model".to_string(),
821        );
822        assert!(
823            llm_request.stream_events.is_some(),
824            "providers that require streaming need a no-op sender even when direct caller did not request one"
825        );
826    }
827}