1use std::sync::Arc;
5
6use bitflags::bitflags;
7use chrono::{DateTime, Utc};
8use serde::{Deserialize, Serialize};
9use serde_json::json;
10use typed_builder::TypedBuilder;
11use uuid::Uuid;
12
13use crate::api::runtime::NemoRelayContextState;
14use crate::api::runtime::current_scope_stack;
15use crate::api::runtime::global_context;
16use crate::api::runtime::{
17 LlmCollectorFn, LlmExecutionNextFn, LlmFinalizerFn, LlmJsonStream, LlmStreamExecutionNextFn,
18};
19use crate::api::scope::event;
20use crate::api::scope::{EmitMarkEventParams, ScopeHandle};
21use crate::api::shared::{
22 ensure_runtime_owner, resolve_parent_uuid, run_request_intercepts_with_codec,
23 snapshot_event_subscribers,
24};
25use crate::codec::request::AnnotatedLlmRequest;
26use crate::codec::response::AnnotatedLlmResponse;
27use crate::codec::traits::{LlmCodec, LlmResponseCodec};
28use crate::error::{FlowError, Result};
29use crate::json::Json;
30use crate::stream::LlmStreamWrapper;
31
32bitflags! {
33 #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
35 pub struct LlmAttributes: u32 {
36 const STATEFUL = 0b01;
38 const STREAMING = 0b10;
40 }
41}
42
43#[derive(Debug, Clone, Serialize, Deserialize, TypedBuilder)]
45#[builder(field_defaults(setter(strip_option(ignore_invalid, fallback_suffix = "_opt"))))]
46pub struct LlmHandle {
47 #[builder(default = Uuid::now_v7())]
49 pub uuid: Uuid,
50 #[builder(default = Utc::now())]
52 pub started_at: DateTime<Utc>,
53 #[builder(setter(into))]
55 pub name: String,
56 #[builder(default)]
58 pub data: Option<Json>,
59 #[builder(default)]
61 pub metadata: Option<Json>,
62 #[builder(default = LlmAttributes::empty())]
64 pub attributes: LlmAttributes,
65 #[builder(default)]
67 pub parent_uuid: Option<Uuid>,
68 #[builder(default, setter(into))]
70 pub model_name: Option<String>,
71}
72
73#[derive(Debug, Clone, Serialize, Deserialize)]
75pub struct LlmRequest {
76 pub headers: serde_json::Map<String, Json>,
78 pub content: Json,
80}
81
82#[derive(Debug, Clone, TypedBuilder)]
84#[builder(field_defaults(setter(strip_option(ignore_invalid, fallback_suffix = "_opt"))))]
85pub struct CreateLlmHandleParams<'a> {
86 pub name: &'a str,
88 #[builder(default)]
90 pub parent_uuid: Option<uuid::Uuid>,
91 #[builder(default = LlmAttributes::empty())]
93 pub attributes: LlmAttributes,
94 #[builder(default)]
96 pub data: Option<Json>,
97 #[builder(default)]
99 pub metadata: Option<Json>,
100 #[builder(default, setter(into))]
102 pub model_name: Option<String>,
103 #[builder(default)]
106 pub timestamp: Option<DateTime<Utc>>,
107}
108
109#[derive(Clone, TypedBuilder)]
111#[builder(field_defaults(setter(strip_option(ignore_invalid, fallback_suffix = "_opt"))))]
112pub struct EndLlmHandleParams<'a> {
113 pub handle: &'a LlmHandle,
115 #[builder(default)]
117 pub data: Option<Json>,
118 #[builder(default)]
120 pub metadata: Option<Json>,
121 #[builder(default)]
123 pub annotated_response: Option<Arc<AnnotatedLlmResponse>>,
124 #[builder(default)]
128 pub timestamp: Option<DateTime<Utc>>,
129}
130
131#[derive(TypedBuilder)]
133#[builder(field_defaults(setter(strip_option(ignore_invalid, fallback_suffix = "_opt"))))]
134pub struct LlmCallParams<'a> {
135 pub name: &'a str,
137 pub request: &'a LlmRequest,
139 #[builder(default)]
141 pub parent: Option<&'a ScopeHandle>,
142 #[builder(default = LlmAttributes::empty())]
144 pub attributes: LlmAttributes,
145 #[builder(default)]
148 pub data: Option<Json>,
149 #[builder(default)]
151 pub metadata: Option<Json>,
152 #[builder(default, setter(into))]
154 pub model_name: Option<String>,
155 #[builder(default)]
157 pub annotated_request: Option<Arc<AnnotatedLlmRequest>>,
158 #[builder(default)]
161 pub timestamp: Option<DateTime<Utc>>,
162}
163
164#[derive(TypedBuilder)]
166#[builder(field_defaults(setter(strip_option(ignore_invalid, fallback_suffix = "_opt"))))]
167pub struct LlmCallExecuteParams {
168 #[builder(setter(into))]
170 pub name: String,
171 pub request: LlmRequest,
173 pub func: LlmExecutionNextFn,
175 #[builder(default)]
177 pub parent: Option<ScopeHandle>,
178 #[builder(default = LlmAttributes::empty())]
180 pub attributes: LlmAttributes,
181 #[builder(default)]
184 pub data: Option<Json>,
185 #[builder(default)]
187 pub metadata: Option<Json>,
188 #[builder(default, setter(into))]
190 pub model_name: Option<String>,
191 #[builder(default)]
193 pub codec: Option<Arc<dyn LlmCodec>>,
194 #[builder(default)]
196 pub response_codec: Option<Arc<dyn LlmResponseCodec>>,
197}
198
199#[derive(TypedBuilder)]
201#[builder(field_defaults(setter(strip_option(ignore_invalid, fallback_suffix = "_opt"))))]
202pub struct LlmStreamCallExecuteParams {
203 #[builder(setter(into))]
205 pub name: String,
206 pub request: LlmRequest,
208 pub func: LlmStreamExecutionNextFn,
210 pub collector: LlmCollectorFn,
212 pub finalizer: LlmFinalizerFn,
214 #[builder(default)]
216 pub parent: Option<ScopeHandle>,
217 #[builder(default = LlmAttributes::empty())]
219 pub attributes: LlmAttributes,
220 #[builder(default)]
223 pub data: Option<Json>,
224 #[builder(default)]
226 pub metadata: Option<Json>,
227 #[builder(default, setter(into))]
229 pub model_name: Option<String>,
230 #[builder(default)]
232 pub codec: Option<Arc<dyn LlmCodec>>,
233 #[builder(default)]
235 pub response_codec: Option<Arc<dyn LlmResponseCodec>>,
236}
237
238#[derive(TypedBuilder)]
240#[builder(field_defaults(setter(strip_option(ignore_invalid, fallback_suffix = "_opt"))))]
241pub struct LlmCallEndParams<'a> {
242 pub handle: &'a LlmHandle,
244 pub response: Json,
246 #[builder(default)]
249 pub data: Option<Json>,
250 #[builder(default)]
252 pub metadata: Option<Json>,
253 #[builder(default)]
255 pub annotated_response: Option<Arc<AnnotatedLlmResponse>>,
256 #[builder(default)]
258 pub response_codec: Option<Arc<dyn LlmResponseCodec>>,
259 #[builder(default)]
263 pub timestamp: Option<DateTime<Utc>>,
264}
265
266fn create_llm_handle(params: CreateLlmHandleParams<'_>) -> Result<LlmHandle> {
267 ensure_runtime_owner()?;
268 let context = global_context();
269 let state = context
270 .read()
271 .map_err(|error| FlowError::Internal(error.to_string()))?;
272 Ok(state.create_llm_handle(params))
273}
274
275fn emit_llm_start(
276 handle: &LlmHandle,
277 request: &LlmRequest,
278 annotated_request: Option<Arc<AnnotatedLlmRequest>>,
279) -> Result<()> {
280 ensure_runtime_owner()?;
281 let (event, subscribers) = {
282 let scope_stack = current_scope_stack();
283 let scope_guard = scope_stack.read().expect("scope stack lock poisoned");
284 let scope_locals = scope_guard.collect_scope_local_registries(|registries| {
285 ®istries.llm_sanitize_request_guardrails
286 });
287 let scope_subscribers = scope_guard.collect_scope_local_subscribers();
288 let subscribers = snapshot_event_subscribers(scope_subscribers)?;
289 let context = global_context();
290 let state = context
291 .read()
292 .map_err(|error| FlowError::Internal(error.to_string()))?;
293
294 let sanitized_request = state.llm_sanitize_request_chain(request.clone(), &scope_locals);
295 let input = serde_json::to_value(&sanitized_request).unwrap_or(Json::Null);
296 let event = state.build_llm_start_event(handle, Some(input), annotated_request);
297 (event, subscribers)
298 };
299 NemoRelayContextState::emit_event(&event, &subscribers);
300 Ok(())
301}
302
303pub fn llm_call(params: LlmCallParams<'_>) -> Result<LlmHandle> {
334 let handle_params = CreateLlmHandleParams::builder()
335 .name(params.name)
336 .parent_uuid_opt(resolve_parent_uuid(params.parent))
337 .attributes(params.attributes)
338 .data_opt(params.data)
339 .metadata_opt(params.metadata)
340 .model_name_opt(params.model_name)
341 .timestamp_opt(params.timestamp)
342 .build();
343 let handle = create_llm_handle(handle_params)?;
344 emit_llm_start(&handle, params.request, params.annotated_request)?;
345 Ok(handle)
346}
347
348pub fn llm_call_end(params: LlmCallEndParams<'_>) -> Result<()> {
380 let LlmCallEndParams {
381 handle,
382 response,
383 data,
384 metadata,
385 annotated_response,
386 response_codec,
387 timestamp,
388 } = params;
389 ensure_runtime_owner()?;
390 let mut decode_error = None;
391 let (event, subscribers) = {
392 let scope_stack = current_scope_stack();
393 let scope_guard = scope_stack.read().expect("scope stack lock poisoned");
394 let scope_locals = scope_guard.collect_scope_local_registries(|registries| {
395 ®istries.llm_sanitize_response_guardrails
396 });
397 let scope_subscribers = scope_guard.collect_scope_local_subscribers();
398 let subscribers = snapshot_event_subscribers(scope_subscribers)?;
399 let context = global_context();
400 let state = context
401 .read()
402 .map_err(|error| FlowError::Internal(error.to_string()))?;
403
404 let sanitized_response = state.llm_sanitize_response_chain(response, &scope_locals);
405 let data = if sanitized_response.is_null() {
406 data
407 } else {
408 Some(sanitized_response)
409 };
410 let annotated_response = match annotated_response {
411 Some(annotated_response) => Some(annotated_response),
412 None => match (response_codec.as_ref(), data.as_ref()) {
413 (Some(codec), Some(response)) => match codec.decode_response(response) {
414 Ok(decoded) => Some(Arc::new(decoded)),
415 Err(error) => {
416 decode_error = Some(error);
417 None
418 }
419 },
420 _ => None,
421 },
422 };
423 let event = state.build_llm_end_event(
424 EndLlmHandleParams::builder()
425 .handle(handle)
426 .data_opt(data)
427 .metadata_opt(metadata)
428 .annotated_response_opt(annotated_response)
429 .timestamp_opt(timestamp)
430 .build(),
431 );
432 (event, subscribers)
433 };
434 NemoRelayContextState::emit_event(&event, &subscribers);
435 if let Some(error) = decode_error {
436 Err(error)
437 } else {
438 Ok(())
439 }
440}
441
442fn emit_llm_end_without_output(handle: &LlmHandle, metadata: Option<Json>) -> Result<()> {
443 ensure_runtime_owner()?;
444 let (event, subscribers) = {
445 let scope_stack = current_scope_stack();
446 let scope_guard = scope_stack.read().expect("scope stack lock poisoned");
447 let scope_subscribers = scope_guard.collect_scope_local_subscribers();
448 let subscribers = snapshot_event_subscribers(scope_subscribers)?;
449 let context = global_context();
450 let state = context
451 .read()
452 .map_err(|error| FlowError::Internal(error.to_string()))?;
453 let event = state.end_llm_handle(handle, handle.data.clone(), metadata, None);
454 (event, subscribers)
455 };
456 NemoRelayContextState::emit_event(&event, &subscribers);
457 Ok(())
458}
459
460pub async fn llm_call_execute(params: LlmCallExecuteParams) -> Result<Json> {
499 let LlmCallExecuteParams {
500 name,
501 request,
502 func,
503 parent,
504 attributes,
505 data,
506 metadata,
507 model_name,
508 codec,
509 response_codec,
510 } = params;
511 ensure_runtime_owner()?;
512 {
513 let (entries, subscribers, parent_uuid, guardrail_metadata) = {
514 let scope_stack = current_scope_stack();
515 let scope_guard = scope_stack.read().expect("scope stack lock poisoned");
516 let scope_locals = scope_guard.collect_scope_local_registries(|registries| {
517 ®istries.llm_conditional_execution_guardrails
518 });
519 let scope_subscribers = scope_guard.collect_scope_local_subscribers();
520 let context = global_context();
521 let state = context
522 .read()
523 .map_err(|error| FlowError::Internal(error.to_string()))?;
524 let entries = state.llm_conditional_execution_entries(&scope_locals);
525 let subscribers = state.collect_event_subscribers(&scope_subscribers);
526 (
527 entries,
528 subscribers,
529 resolve_parent_uuid(parent.as_ref()),
530 metadata.clone(),
531 )
532 };
533 if let Some(error) = NemoRelayContextState::llm_conditional_execution_snapshot_chain(
534 &request,
535 &entries,
536 &subscribers,
537 parent_uuid,
538 guardrail_metadata,
539 )? {
540 let mut rejection_data = json!({});
541 if let Some(object) = rejection_data.as_object_mut() {
542 object.insert("rejected".into(), json!(true));
543 object.insert("rejection_reason".into(), json!(&error));
544 }
545 let _ = event(
546 EmitMarkEventParams::builder()
547 .name(&name)
548 .parent_opt(parent.as_ref())
549 .data(rejection_data)
550 .metadata_opt(metadata.clone())
551 .build(),
552 );
553 return Err(FlowError::GuardrailRejected(error));
554 }
555 }
556
557 let (intercepted_request, annotated_request) =
558 run_request_intercepts_with_codec(&name, request, codec)?;
559
560 let handle = create_llm_handle(
561 CreateLlmHandleParams::builder()
562 .name(name.as_str())
563 .parent_uuid_opt(resolve_parent_uuid(parent.as_ref()))
564 .attributes(attributes)
565 .data_opt(data.clone())
566 .metadata_opt(metadata.clone())
567 .model_name_opt(model_name)
568 .build(),
569 )?;
570 emit_llm_start(&handle, &intercepted_request, annotated_request.clone())?;
571
572 let execution = {
573 let scope_stack = current_scope_stack();
574 let scope_guard = scope_stack.read().expect("scope stack lock poisoned");
575 let scope_locals = scope_guard
576 .collect_scope_local_registries(|registries| ®istries.llm_execution_intercepts);
577 let context = global_context();
578 let state = context
579 .read()
580 .map_err(|error| FlowError::Internal(error.to_string()))?;
581 state.llm_build_execution_chain(&name, func, &scope_locals)
582 };
583
584 match execution(intercepted_request).await {
585 Ok(response) => {
586 let annotated_response = response_codec
587 .as_ref()
588 .and_then(|codec| codec.decode_response(&response).ok())
589 .map(Arc::new);
590 llm_call_end(
591 LlmCallEndParams::builder()
592 .handle(&handle)
593 .response(response.clone())
594 .data_opt(data)
595 .metadata_opt(metadata)
596 .annotated_response_opt(annotated_response)
597 .build(),
598 )?;
599 Ok(response)
600 }
601 Err(error) => {
602 let _ = emit_llm_end_without_output(&handle, metadata);
603 Err(error)
604 }
605 }
606}
607
608pub async fn llm_stream_call_execute(params: LlmStreamCallExecuteParams) -> Result<LlmJsonStream> {
645 let LlmStreamCallExecuteParams {
646 name,
647 request,
648 func,
649 collector,
650 finalizer,
651 parent,
652 attributes,
653 data,
654 metadata,
655 model_name,
656 codec,
657 response_codec,
658 } = params;
659 ensure_runtime_owner()?;
660 {
661 let (entries, subscribers, parent_uuid, guardrail_metadata) = {
662 let scope_stack = current_scope_stack();
663 let scope_guard = scope_stack.read().expect("scope stack lock poisoned");
664 let scope_locals = scope_guard.collect_scope_local_registries(|registries| {
665 ®istries.llm_conditional_execution_guardrails
666 });
667 let scope_subscribers = scope_guard.collect_scope_local_subscribers();
668 let context = global_context();
669 let state = context
670 .read()
671 .map_err(|error| FlowError::Internal(error.to_string()))?;
672 let entries = state.llm_conditional_execution_entries(&scope_locals);
673 let subscribers = state.collect_event_subscribers(&scope_subscribers);
674 (
675 entries,
676 subscribers,
677 resolve_parent_uuid(parent.as_ref()),
678 metadata.clone(),
679 )
680 };
681 if let Some(error) = NemoRelayContextState::llm_conditional_execution_snapshot_chain(
682 &request,
683 &entries,
684 &subscribers,
685 parent_uuid,
686 guardrail_metadata,
687 )? {
688 let mut rejection_data = json!({});
689 if let Some(object) = rejection_data.as_object_mut() {
690 object.insert("rejected".into(), json!(true));
691 object.insert("rejection_reason".into(), json!(&error));
692 }
693 let _ = event(
694 EmitMarkEventParams::builder()
695 .name(&name)
696 .parent_opt(parent.as_ref())
697 .data(rejection_data)
698 .metadata_opt(metadata.clone())
699 .build(),
700 );
701 return Err(FlowError::GuardrailRejected(error));
702 }
703 }
704
705 let (intercepted_request, annotated_request) =
706 run_request_intercepts_with_codec(&name, request, codec)?;
707
708 let handle = create_llm_handle(
709 CreateLlmHandleParams::builder()
710 .name(name.as_str())
711 .parent_uuid_opt(resolve_parent_uuid(parent.as_ref()))
712 .attributes(attributes)
713 .data_opt(data.clone())
714 .metadata_opt(metadata.clone())
715 .model_name_opt(model_name)
716 .build(),
717 )?;
718 emit_llm_start(&handle, &intercepted_request, annotated_request)?;
719
720 let execution = {
721 let scope_stack = current_scope_stack();
722 let scope_guard = scope_stack.read().expect("scope stack lock poisoned");
723 let scope_locals = scope_guard.collect_scope_local_registries(|registries| {
724 ®istries.llm_stream_execution_intercepts
725 });
726 let context = global_context();
727 let state = context
728 .read()
729 .map_err(|error| FlowError::Internal(error.to_string()))?;
730 state.llm_stream_build_execution_chain(&name, func, &scope_locals)
731 };
732
733 match execution(intercepted_request).await {
734 Ok(raw_stream) => {
735 let wrapper = LlmStreamWrapper::new(
736 raw_stream,
737 handle,
738 collector,
739 finalizer,
740 data,
741 metadata,
742 response_codec,
743 );
744 Ok(Box::pin(wrapper) as LlmJsonStream)
745 }
746 Err(error) => {
747 let _ = emit_llm_end_without_output(&handle, metadata);
748 Err(error)
749 }
750 }
751}
752
753pub fn llm_request_intercepts(name: &str, request: LlmRequest) -> Result<LlmRequest> {
773 ensure_runtime_owner()?;
774 let scope_stack = current_scope_stack();
775 let scope_guard = scope_stack.read().expect("scope stack lock poisoned");
776 let scope_locals =
777 scope_guard.collect_scope_local_registries(|registries| ®istries.llm_request_intercepts);
778 let context = global_context();
779 let state = context
780 .read()
781 .map_err(|error| FlowError::Internal(error.to_string()))?;
782 let (request, _) = state.llm_request_intercepts_chain(name, request, None, &scope_locals)?;
783 Ok(request)
784}
785
786pub fn llm_conditional_execution(request: &LlmRequest) -> Result<()> {
807 ensure_runtime_owner()?;
808 let (entries, subscribers, parent_uuid) = {
809 let scope_stack = current_scope_stack();
810 let scope_guard = scope_stack.read().expect("scope stack lock poisoned");
811 let scope_locals = scope_guard.collect_scope_local_registries(|registries| {
812 ®istries.llm_conditional_execution_guardrails
813 });
814 let scope_subscribers = scope_guard.collect_scope_local_subscribers();
815 let context = global_context();
816 let state = context
817 .read()
818 .map_err(|error| FlowError::Internal(error.to_string()))?;
819 let entries = state.llm_conditional_execution_entries(&scope_locals);
820 let subscribers = state.collect_event_subscribers(&scope_subscribers);
821 (entries, subscribers, resolve_parent_uuid(None))
822 };
823 if let Some(error) = NemoRelayContextState::llm_conditional_execution_snapshot_chain(
824 request,
825 &entries,
826 &subscribers,
827 parent_uuid,
828 None,
829 )? {
830 return Err(FlowError::GuardrailRejected(error));
831 }
832 Ok(())
833}