///|
/// A provider-neutral request built from mizchi/llm message and tool types.
pub(all) struct LlmRequest {
messages : Array[@llm.Message]
tools : Array[@llm.ToolDef]
}
///|
/// Creates one LLM invocation request.
pub fn LlmRequest::LlmRequest(
messages : Array[@llm.Message],
tools? : Array[@llm.ToolDef] = [],
) -> LlmRequest {
LlmRequest::{ messages, tools }
}
///|
pub struct LlmNodeSpec[S, P] {
provider : @llm.BoxedProvider
build_request : (@core.NodeContext, S) -> LlmRequest raise
decode_response : (S, @llm.CollectResult) -> @core.NodeOutput[P] raise
}
///|
/// Raised when an LLM provider reports a stream failure.
pub suberror LlmNodeError {
ProviderFailed(String)
} derive(Debug)
///|
fn collect_response(
provider : @llm.BoxedProvider,
request : LlmRequest,
) -> @llm.CollectResult raise LlmNodeError {
let text = StringBuilder::new()
let tool_calls : Array[@llm.ToolCall] = []
let finish_reason : Ref[@llm.FinishReason] = Ref(@llm.Stop)
let usage : Ref[@llm.Usage?] = Ref(None)
let failure : Ref[String?] = Ref(None)
provider.stream(request.messages, request.tools, @llm.StreamHandler::{
on_event: event => {
match event {
@llm.TextDelta(value) => text.write_string(value)
@llm.ToolCallEnd(id~, name~, input~) =>
tool_calls.push(@llm.ToolCall::{ id, name, input })
@llm.MessageEnd(finish_reason=reason, usage=reported_usage) => {
finish_reason.val = reason
usage.val = reported_usage
}
@llm.Error(message) => failure.val = Some(message)
_ => ()
}
},
})
match failure.val {
Some(message) => raise ProviderFailed(message)
None =>
@llm.CollectResult::{
text: text.to_string(),
tool_calls,
finish_reason: finish_reason.val,
usage: usage.val,
}
}
}
///|
/// Creates a provider-backed LLM graph-node specification.
pub fn[S, P] LlmNodeSpec::LlmNodeSpec(
provider : @llm.BoxedProvider,
build_request : (@core.NodeContext, S) -> LlmRequest raise,
decode_response : (S, @llm.CollectResult) -> @core.NodeOutput[P] raise,
) -> LlmNodeSpec[S, P] {
LlmNodeSpec::{ provider, build_request, decode_response }
}
///|
/// Creates a graph node that builds a request, invokes its provider, and
/// decodes the collected response.
pub fn[S, P] llm_node(
id : @core.NodeId,
metadata : @core.NodeMetadata,
spec : LlmNodeSpec[S, P],
) -> @core.Node[S, P] {
@core.Node::Node(id, metadata, async fn(context, state) {
let request = (spec.build_request)(context, state)
@async.pause()
let response = collect_response(spec.provider, request)
(spec.decode_response)(state, response)
})
}