mirror of
https://github.com/toeverything/AFFiNE.git
synced 2026-09-02 06:39:46 +08:00
feat(server): refactor copilot (#14892)
#### PR Dependency Tree * **PR #14892** 👈 This tree was auto-generated by [Charcoal](https://github.com/danerwilliams/charcoal)
This commit is contained in:
@@ -0,0 +1,13 @@
|
||||
use napi::{Error, Status};
|
||||
|
||||
pub(crate) const STREAM_END_MARKER: &str = "__AFFINE_LLM_STREAM_END__";
|
||||
pub(crate) const STREAM_ABORTED_REASON: &str = "__AFFINE_LLM_STREAM_ABORTED__";
|
||||
pub(crate) const STREAM_CALLBACK_DISPATCH_FAILED_REASON: &str = "__AFFINE_LLM_STREAM_CALLBACK_DISPATCH_FAILED__";
|
||||
|
||||
pub(crate) fn callback_dispatch_failed_reason(status: Status) -> String {
|
||||
format!("{STREAM_CALLBACK_DISPATCH_FAILED_REASON}:{status}")
|
||||
}
|
||||
|
||||
pub(crate) fn invalid_arg(message: impl Into<String>) -> Error {
|
||||
Error::new(Status::InvalidArg, message.into())
|
||||
}
|
||||
@@ -0,0 +1,15 @@
|
||||
mod error;
|
||||
mod stream;
|
||||
mod stream_handle;
|
||||
mod tool_loop;
|
||||
|
||||
pub(crate) use error::{
|
||||
STREAM_ABORTED_REASON, STREAM_CALLBACK_DISPATCH_FAILED_REASON, STREAM_END_MARKER, callback_dispatch_failed_reason,
|
||||
invalid_arg,
|
||||
};
|
||||
pub(crate) use stream::emit_error_event;
|
||||
pub use stream::{
|
||||
llm_dispatch_prepared_stream, llm_dispatch_tool_loop_stream, llm_dispatch_tool_loop_stream_prepared,
|
||||
llm_dispatch_tool_loop_stream_routed,
|
||||
};
|
||||
pub(crate) use stream_handle::LlmStreamHandle;
|
||||
@@ -0,0 +1,230 @@
|
||||
use std::sync::{
|
||||
Arc,
|
||||
atomic::{AtomicBool, Ordering},
|
||||
};
|
||||
|
||||
use llm_adapter::{
|
||||
backend::{BackendConfig, BackendError, BackendHttpClient, DefaultHttpClient},
|
||||
core::StreamEvent,
|
||||
router::{PreparedChatRoute, RoutedBackend, dispatch_prepared_stream_with_pipeline},
|
||||
};
|
||||
use napi::{
|
||||
Result, Status,
|
||||
bindgen_prelude::PromiseRaw,
|
||||
threadsafe_function::{ThreadsafeFunction, ThreadsafeFunctionCallMode},
|
||||
};
|
||||
|
||||
use super::{STREAM_CALLBACK_DISPATCH_FAILED_REASON, STREAM_END_MARKER, callback_dispatch_failed_reason, tool_loop};
|
||||
use crate::llm::{
|
||||
LlmDispatchPayload, LlmRoutedBackendPayload, LlmStreamHandle, STREAM_ABORTED_REASON, StreamPipeline,
|
||||
backend_transport_error, map_json_error, parse_prepared_chat_routes_with_middleware,
|
||||
parse_prepared_chat_routes_without_middleware, parse_protocol, resolve_stream_chain,
|
||||
};
|
||||
|
||||
type PreparedDispatchRoute = (PreparedChatRoute, crate::llm::LlmMiddlewarePayload);
|
||||
|
||||
#[napi(catch_unwind)]
|
||||
pub fn llm_dispatch_prepared_stream(
|
||||
routes_json: String,
|
||||
callback: ThreadsafeFunction<String, ()>,
|
||||
) -> Result<LlmStreamHandle> {
|
||||
let routes = parse_prepared_chat_routes_with_middleware(&routes_json)?;
|
||||
Ok(spawn_prepared_stream(routes, callback))
|
||||
}
|
||||
|
||||
#[napi(catch_unwind)]
|
||||
pub fn llm_dispatch_tool_loop_stream(
|
||||
protocol: String,
|
||||
backend_config_json: String,
|
||||
request_json: String,
|
||||
max_steps: u32,
|
||||
callback: ThreadsafeFunction<String, ()>,
|
||||
tool_callback: ThreadsafeFunction<String, PromiseRaw<'static, String>>,
|
||||
) -> Result<LlmStreamHandle> {
|
||||
let protocol = parse_protocol(&protocol)?;
|
||||
let config: BackendConfig = serde_json::from_str(&backend_config_json).map_err(map_json_error)?;
|
||||
let payload: LlmDispatchPayload = serde_json::from_str(&request_json).map_err(map_json_error)?;
|
||||
|
||||
Ok(tool_loop::spawn_tool_loop_stream(
|
||||
protocol,
|
||||
config,
|
||||
payload,
|
||||
max_steps as usize,
|
||||
callback,
|
||||
tool_callback,
|
||||
))
|
||||
}
|
||||
|
||||
#[napi(catch_unwind)]
|
||||
pub fn llm_dispatch_tool_loop_stream_routed(
|
||||
routes_json: String,
|
||||
request_json: String,
|
||||
max_steps: u32,
|
||||
callback: ThreadsafeFunction<String, ()>,
|
||||
tool_callback: ThreadsafeFunction<String, PromiseRaw<'static, String>>,
|
||||
) -> Result<LlmStreamHandle> {
|
||||
let routes = parse_routed_backends(&routes_json)?;
|
||||
let payload: LlmDispatchPayload = serde_json::from_str(&request_json).map_err(map_json_error)?;
|
||||
|
||||
Ok(tool_loop::spawn_routed_tool_loop_stream(
|
||||
routes,
|
||||
payload,
|
||||
max_steps as usize,
|
||||
callback,
|
||||
tool_callback,
|
||||
))
|
||||
}
|
||||
|
||||
#[napi(catch_unwind)]
|
||||
pub fn llm_dispatch_tool_loop_stream_prepared(
|
||||
routes_json: String,
|
||||
max_steps: u32,
|
||||
callback: ThreadsafeFunction<String, ()>,
|
||||
tool_callback: ThreadsafeFunction<String, PromiseRaw<'static, String>>,
|
||||
) -> Result<LlmStreamHandle> {
|
||||
let routes = parse_prepared_chat_routes_without_middleware(&routes_json)?;
|
||||
Ok(tool_loop::spawn_prepared_tool_loop_stream(
|
||||
routes,
|
||||
max_steps as usize,
|
||||
callback,
|
||||
tool_callback,
|
||||
))
|
||||
}
|
||||
|
||||
fn spawn_prepared_stream(
|
||||
routes: Vec<PreparedDispatchRoute>,
|
||||
callback: ThreadsafeFunction<String, ()>,
|
||||
) -> LlmStreamHandle {
|
||||
let aborted = Arc::new(AtomicBool::new(false));
|
||||
let aborted_in_worker = aborted.clone();
|
||||
|
||||
std::thread::spawn(move || {
|
||||
let result = dispatch_prepared_stream_with_fallback(&routes, &callback, &aborted_in_worker);
|
||||
let callback_dispatch_failed = matches!(
|
||||
&result,
|
||||
Err(BackendError::Transport { message: reason })
|
||||
if reason.starts_with(STREAM_CALLBACK_DISPATCH_FAILED_REASON)
|
||||
);
|
||||
|
||||
if let Err(error) = result
|
||||
&& !aborted_in_worker.load(Ordering::Relaxed)
|
||||
&& !callback_dispatch_failed
|
||||
&& !is_abort_error(&error)
|
||||
{
|
||||
emit_error_event(&callback, error.to_string(), "dispatch_error");
|
||||
}
|
||||
|
||||
if !callback_dispatch_failed {
|
||||
let _ = callback.call(
|
||||
Ok(STREAM_END_MARKER.to_string()),
|
||||
ThreadsafeFunctionCallMode::NonBlocking,
|
||||
);
|
||||
}
|
||||
});
|
||||
|
||||
LlmStreamHandle { aborted }
|
||||
}
|
||||
|
||||
fn dispatch_prepared_stream_with_fallback(
|
||||
routes: &[PreparedDispatchRoute],
|
||||
callback: &ThreadsafeFunction<String, ()>,
|
||||
aborted: &AtomicBool,
|
||||
) -> std::result::Result<(), BackendError> {
|
||||
dispatch_prepared_stream_with_fallback_using_client(&DefaultHttpClient::default(), routes, aborted, |event| {
|
||||
emit_stream_event(callback, event)
|
||||
})
|
||||
}
|
||||
|
||||
fn dispatch_prepared_stream_with_fallback_using_client<F>(
|
||||
client: &dyn BackendHttpClient,
|
||||
routes: &[PreparedDispatchRoute],
|
||||
aborted: &AtomicBool,
|
||||
mut emit_event: F,
|
||||
) -> std::result::Result<(), BackendError>
|
||||
where
|
||||
F: FnMut(&StreamEvent) -> Status,
|
||||
{
|
||||
let mut adapter_routes = routes
|
||||
.iter()
|
||||
.map(|(route, middleware)| {
|
||||
let chain =
|
||||
resolve_stream_chain(&middleware.stream).map_err(|error| backend_transport_error(error.reason.clone()))?;
|
||||
Ok((route.clone(), StreamPipeline::new(chain, middleware.config.clone())))
|
||||
})
|
||||
.collect::<std::result::Result<Vec<_>, BackendError>>()?;
|
||||
let mut callback_dispatch_failed = false;
|
||||
|
||||
dispatch_prepared_stream_with_pipeline(
|
||||
client,
|
||||
&mut adapter_routes,
|
||||
|| aborted.load(Ordering::Relaxed),
|
||||
|| backend_transport_error(STREAM_ABORTED_REASON),
|
||||
|event| {
|
||||
let status = emit_event(event);
|
||||
if status != Status::Ok {
|
||||
callback_dispatch_failed = true;
|
||||
return Err(backend_transport_error(callback_dispatch_failed_reason(status)));
|
||||
}
|
||||
Ok(())
|
||||
},
|
||||
)?;
|
||||
|
||||
if callback_dispatch_failed {
|
||||
Err(backend_transport_error(format!(
|
||||
"{STREAM_CALLBACK_DISPATCH_FAILED_REASON}:unknown"
|
||||
)))
|
||||
} else {
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn emit_error_event(callback: &ThreadsafeFunction<String, ()>, message: String, code: &str) {
|
||||
let error_event = serde_json::to_string(&StreamEvent::Error {
|
||||
message: message.clone(),
|
||||
code: Some(code.to_string()),
|
||||
})
|
||||
.unwrap_or_else(|_| {
|
||||
serde_json::json!({
|
||||
"type": "error",
|
||||
"message": message,
|
||||
"code": code,
|
||||
})
|
||||
.to_string()
|
||||
});
|
||||
|
||||
let _ = callback.call(Ok(error_event), ThreadsafeFunctionCallMode::NonBlocking);
|
||||
}
|
||||
|
||||
fn emit_stream_event(callback: &ThreadsafeFunction<String, ()>, event: &StreamEvent) -> Status {
|
||||
let value = serde_json::to_string(event).unwrap_or_else(|error| {
|
||||
serde_json::json!({
|
||||
"type": "error",
|
||||
"message": format!("failed to serialize stream event: {error}"),
|
||||
})
|
||||
.to_string()
|
||||
});
|
||||
|
||||
callback.call(Ok(value), ThreadsafeFunctionCallMode::NonBlocking)
|
||||
}
|
||||
|
||||
fn parse_routed_backends(routes_json: &str) -> Result<Vec<RoutedBackend>> {
|
||||
let payload: Vec<LlmRoutedBackendPayload> = serde_json::from_str(routes_json).map_err(map_json_error)?;
|
||||
payload
|
||||
.into_iter()
|
||||
.map(|route| {
|
||||
Ok(RoutedBackend {
|
||||
provider_id: route.provider_id,
|
||||
protocol: parse_protocol(&route.protocol)?,
|
||||
model: route.model,
|
||||
config: route.config,
|
||||
})
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn is_abort_error(error: &BackendError) -> bool {
|
||||
matches!(
|
||||
error,
|
||||
BackendError::Transport { message: reason } if reason == STREAM_ABORTED_REASON
|
||||
)
|
||||
}
|
||||
@@ -0,0 +1,17 @@
|
||||
use std::sync::{
|
||||
Arc,
|
||||
atomic::{AtomicBool, Ordering},
|
||||
};
|
||||
|
||||
#[napi]
|
||||
pub struct LlmStreamHandle {
|
||||
pub(crate) aborted: Arc<AtomicBool>,
|
||||
}
|
||||
|
||||
#[napi]
|
||||
impl LlmStreamHandle {
|
||||
#[napi]
|
||||
pub fn abort(&self) {
|
||||
self.aborted.store(true, Ordering::SeqCst);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,165 @@
|
||||
use std::sync::{
|
||||
Arc, Mutex,
|
||||
atomic::{AtomicBool, Ordering},
|
||||
mpsc::{self, SyncSender},
|
||||
};
|
||||
|
||||
use llm_adapter::backend::BackendError;
|
||||
use llm_runtime::{
|
||||
EventSink, ToolCallbackRequest as RuntimeToolCallbackRequest, ToolCallbackResponse as RuntimeToolCallbackResponse,
|
||||
ToolExecutionResult, ToolExecutor, ToolLoopEvent,
|
||||
};
|
||||
use napi::{
|
||||
Error, JsValue, Result, Status,
|
||||
bindgen_prelude::{CallbackContext, PromiseRaw, Unknown},
|
||||
threadsafe_function::{ThreadsafeFunction, ThreadsafeFunctionCallMode},
|
||||
};
|
||||
|
||||
use super::contract::{NativeToolCall, ToolLoopStreamEvent};
|
||||
use crate::llm::{backend_transport_error, host::callback_dispatch_failed_reason};
|
||||
|
||||
type ToolCallbackResult = std::result::Result<RuntimeToolCallbackResponse, String>;
|
||||
type ToolCallbackSender = SyncSender<ToolCallbackResult>;
|
||||
type ToolCallbackSenderSlot = Arc<Mutex<Option<ToolCallbackSender>>>;
|
||||
|
||||
pub(super) struct NapiToolExecutor<'a> {
|
||||
callback: &'a ThreadsafeFunction<String, PromiseRaw<'static, String>>,
|
||||
}
|
||||
|
||||
impl<'a> NapiToolExecutor<'a> {
|
||||
pub(super) fn new(callback: &'a ThreadsafeFunction<String, PromiseRaw<'static, String>>) -> Self {
|
||||
Self { callback }
|
||||
}
|
||||
}
|
||||
|
||||
impl ToolExecutor<BackendError> for NapiToolExecutor<'_> {
|
||||
fn execute(&mut self, call: &NativeToolCall) -> std::result::Result<ToolExecutionResult, BackendError> {
|
||||
let result =
|
||||
execute_tool_callback(self.callback, call).map_err(|error| backend_transport_error(error.to_string()))?;
|
||||
Ok(ToolExecutionResult {
|
||||
call_id: result.call_id,
|
||||
name: result.name,
|
||||
arguments: result.args,
|
||||
arguments_text: result.raw_arguments_text,
|
||||
arguments_error: result.argument_parse_error,
|
||||
output: result.output,
|
||||
is_error: result.is_error,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) struct NapiEventSink<'a> {
|
||||
callback: &'a ThreadsafeFunction<String, ()>,
|
||||
emitted: Option<&'a AtomicBool>,
|
||||
}
|
||||
|
||||
impl<'a> NapiEventSink<'a> {
|
||||
pub(super) fn new_with_emitted(callback: &'a ThreadsafeFunction<String, ()>, emitted: &'a AtomicBool) -> Self {
|
||||
Self {
|
||||
callback,
|
||||
emitted: Some(emitted),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl EventSink<BackendError> for NapiEventSink<'_> {
|
||||
fn emit(&mut self, event: &ToolLoopEvent) -> std::result::Result<(), BackendError> {
|
||||
if let Some(emitted) = self.emitted {
|
||||
emitted.store(true, Ordering::Relaxed);
|
||||
}
|
||||
emit_tool_loop_event(self.callback, event)
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn emit_tool_loop_event(
|
||||
callback: &ThreadsafeFunction<String, ()>,
|
||||
event: &ToolLoopStreamEvent,
|
||||
) -> std::result::Result<(), BackendError> {
|
||||
let value = serde_json::to_string(event).unwrap_or_else(|error| {
|
||||
serde_json::json!({
|
||||
"type": "error",
|
||||
"message": format!("failed to serialize tool loop event: {error}"),
|
||||
})
|
||||
.to_string()
|
||||
});
|
||||
|
||||
let status = callback.call(Ok(value), ThreadsafeFunctionCallMode::NonBlocking);
|
||||
if status != Status::Ok {
|
||||
return Err(backend_transport_error(callback_dispatch_failed_reason(status)));
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub(super) fn execute_tool_callback(
|
||||
callback: &ThreadsafeFunction<String, PromiseRaw<'static, String>>,
|
||||
call: &NativeToolCall,
|
||||
) -> Result<RuntimeToolCallbackResponse> {
|
||||
let request = RuntimeToolCallbackRequest {
|
||||
call_id: call.id.clone(),
|
||||
name: call.name.clone(),
|
||||
args: call.args.clone(),
|
||||
raw_arguments_text: call.raw_arguments_text.clone(),
|
||||
argument_parse_error: call.argument_parse_error.clone(),
|
||||
};
|
||||
let request = serde_json::to_string(&request).map_err(|error| Error::new(Status::InvalidArg, error.to_string()))?;
|
||||
let (sender, receiver) = mpsc::sync_channel::<ToolCallbackResult>(1);
|
||||
let sender = Arc::new(Mutex::new(Some(sender)));
|
||||
let sender_in_callback = sender.clone();
|
||||
let status = callback.call_with_return_value(
|
||||
Ok(request),
|
||||
ThreadsafeFunctionCallMode::NonBlocking,
|
||||
move |promise, _env| {
|
||||
match promise {
|
||||
Ok(promise) => {
|
||||
let sender_in_then = sender_in_callback.clone();
|
||||
let sender_in_catch = sender_in_callback.clone();
|
||||
promise
|
||||
.then(move |ctx| {
|
||||
let result = serde_json::from_str(&ctx.value).map_err(|error| error.to_string());
|
||||
send_tool_callback_result(&sender_in_then, result);
|
||||
Ok(())
|
||||
})?
|
||||
.catch(move |ctx: CallbackContext<Unknown>| {
|
||||
let message = ctx.value.coerce_to_string()?.into_utf8()?.as_str()?.to_string();
|
||||
send_tool_callback_result(&sender_in_catch, Err(message));
|
||||
Ok(())
|
||||
})?;
|
||||
}
|
||||
Err(error) => {
|
||||
send_tool_callback_result(&sender_in_callback, Err(error.to_string()));
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
},
|
||||
);
|
||||
|
||||
if status != Status::Ok {
|
||||
return Err(Error::new(
|
||||
Status::GenericFailure,
|
||||
format!("native tool callback dispatch failed: {status}"),
|
||||
));
|
||||
}
|
||||
|
||||
let response_json = receiver.recv().map_err(|_| {
|
||||
Error::new(
|
||||
Status::GenericFailure,
|
||||
"native tool callback receiver closed before completion",
|
||||
)
|
||||
})?;
|
||||
|
||||
let response = response_json.map_err(|message| Error::new(Status::GenericFailure, message))?;
|
||||
if !response.args.is_object() {
|
||||
return Err(Error::new(
|
||||
Status::InvalidArg,
|
||||
"Tool callback response args must be a JSON object",
|
||||
));
|
||||
}
|
||||
Ok(response)
|
||||
}
|
||||
|
||||
fn send_tool_callback_result(sender: &ToolCallbackSenderSlot, result: ToolCallbackResult) {
|
||||
if let Some(sender) = sender.lock().expect("tool callback sender poisoned").take() {
|
||||
let _ = sender.send(result);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,4 @@
|
||||
use llm_runtime::{AccumulatedToolCall, ToolLoopEvent};
|
||||
|
||||
pub(super) type NativeToolCall = AccumulatedToolCall;
|
||||
pub(super) type ToolLoopStreamEvent = ToolLoopEvent;
|
||||
@@ -0,0 +1,362 @@
|
||||
use std::sync::{
|
||||
Arc,
|
||||
atomic::{AtomicBool, Ordering},
|
||||
};
|
||||
|
||||
use llm_adapter::{
|
||||
backend::{BackendConfig, BackendError, ChatProtocol, DefaultHttpClient},
|
||||
core::CoreRequest,
|
||||
router::{PreparedChatRoute, RoutedBackend, dispatch_prepared_stream_with_fallback_index},
|
||||
};
|
||||
use llm_runtime::{RoundOutcome, RoundProcessorError, run_prepared_stream_round_with_fallback, run_tool_loop};
|
||||
use napi::{
|
||||
bindgen_prelude::PromiseRaw,
|
||||
threadsafe_function::{ThreadsafeFunction, ThreadsafeFunctionCallMode},
|
||||
};
|
||||
|
||||
use super::callback::{NapiEventSink, NapiToolExecutor, emit_tool_loop_event};
|
||||
use crate::llm::{
|
||||
LlmDispatchPayload, LlmMiddlewarePayload, LlmStreamHandle, STREAM_ABORTED_REASON,
|
||||
STREAM_CALLBACK_DISPATCH_FAILED_REASON, STREAM_END_MARKER, StreamPipeline, apply_request_middlewares,
|
||||
backend_transport_error, emit_error_event, resolve_stream_chain,
|
||||
};
|
||||
|
||||
pub(crate) type PreparedToolLoopRoute = (PreparedChatRoute, LlmMiddlewarePayload);
|
||||
|
||||
fn dispatch_prepared_round_with_fallback(
|
||||
routes: &[PreparedToolLoopRoute],
|
||||
callback: &ThreadsafeFunction<String, ()>,
|
||||
aborted: &AtomicBool,
|
||||
emitted: &AtomicBool,
|
||||
) -> std::result::Result<RoundOutcome, BackendError> {
|
||||
let adapter_routes = routes.iter().map(|(route, _)| route.clone()).collect::<Vec<_>>();
|
||||
let mut pipelines = routes
|
||||
.iter()
|
||||
.map(|(_, middleware)| {
|
||||
let chain =
|
||||
resolve_stream_chain(&middleware.stream).map_err(|error| backend_transport_error(error.reason.clone()))?;
|
||||
Ok(StreamPipeline::new(chain, middleware.config.clone()))
|
||||
})
|
||||
.collect::<std::result::Result<Vec<_>, BackendError>>()?;
|
||||
|
||||
run_prepared_stream_round_with_fallback(
|
||||
&mut pipelines,
|
||||
|on_event| {
|
||||
let (selected_index, _) =
|
||||
dispatch_prepared_stream_with_fallback_index(&DefaultHttpClient::default(), &adapter_routes, on_event)?;
|
||||
Ok(selected_index)
|
||||
},
|
||||
|| aborted.load(Ordering::Relaxed),
|
||||
|| backend_transport_error(STREAM_ABORTED_REASON),
|
||||
|error: RoundProcessorError| backend_transport_error(error.to_string()),
|
||||
|loop_event| {
|
||||
emitted.store(true, Ordering::Relaxed);
|
||||
emit_tool_loop_event(callback, loop_event)
|
||||
},
|
||||
)
|
||||
}
|
||||
|
||||
fn prepare_tool_loop_route(
|
||||
route: &RoutedBackend,
|
||||
request: &CoreRequest,
|
||||
middleware: &LlmMiddlewarePayload,
|
||||
) -> std::result::Result<PreparedToolLoopRoute, BackendError> {
|
||||
let mut routed_request =
|
||||
apply_request_middlewares(request.clone(), middleware, route.protocol, route.config.request_layer)
|
||||
.map_err(|error| backend_transport_error(error.reason.clone()))?;
|
||||
routed_request.model = route.model.clone();
|
||||
|
||||
Ok(((route.clone(), routed_request), middleware.clone()))
|
||||
}
|
||||
|
||||
fn dispatch_round(
|
||||
route: &RoutedBackend,
|
||||
request: &CoreRequest,
|
||||
callback: &ThreadsafeFunction<String, ()>,
|
||||
middleware: &LlmMiddlewarePayload,
|
||||
aborted: &AtomicBool,
|
||||
emitted: &AtomicBool,
|
||||
) -> std::result::Result<RoundOutcome, BackendError> {
|
||||
let prepared = vec![prepare_tool_loop_route(route, request, middleware)?];
|
||||
dispatch_prepared_round_with_fallback(&prepared, callback, aborted, emitted)
|
||||
}
|
||||
|
||||
fn dispatch_round_with_fallback(
|
||||
routes: &[RoutedBackend],
|
||||
request: &CoreRequest,
|
||||
callback: &ThreadsafeFunction<String, ()>,
|
||||
middleware: &LlmMiddlewarePayload,
|
||||
aborted: &AtomicBool,
|
||||
emitted: &AtomicBool,
|
||||
) -> std::result::Result<RoundOutcome, BackendError> {
|
||||
let prepared = routes
|
||||
.iter()
|
||||
.map(|route| prepare_tool_loop_route(route, request, middleware))
|
||||
.collect::<std::result::Result<Vec<_>, BackendError>>()?;
|
||||
|
||||
dispatch_prepared_round_with_fallback(&prepared, callback, aborted, emitted)
|
||||
}
|
||||
|
||||
fn dispatch_prepared_payload_round_with_fallback(
|
||||
routes: &[PreparedToolLoopRoute],
|
||||
request: &CoreRequest,
|
||||
callback: &ThreadsafeFunction<String, ()>,
|
||||
aborted: &AtomicBool,
|
||||
emitted: &AtomicBool,
|
||||
) -> std::result::Result<RoundOutcome, BackendError> {
|
||||
let prepared = routes
|
||||
.iter()
|
||||
.map(|((route, _), middleware)| prepare_tool_loop_route(route, request, middleware))
|
||||
.collect::<std::result::Result<Vec<_>, BackendError>>()?;
|
||||
|
||||
dispatch_prepared_round_with_fallback(&prepared, callback, aborted, emitted)
|
||||
}
|
||||
|
||||
fn run_native_tool_loop_with_dispatch<F>(
|
||||
payload: LlmDispatchPayload,
|
||||
max_steps: usize,
|
||||
callback: &ThreadsafeFunction<String, ()>,
|
||||
tool_callback: &ThreadsafeFunction<String, PromiseRaw<'static, String>>,
|
||||
aborted: Arc<AtomicBool>,
|
||||
emitted: &AtomicBool,
|
||||
dispatch_round_fn: F,
|
||||
) -> std::result::Result<(), BackendError>
|
||||
where
|
||||
F: Fn(
|
||||
&CoreRequest,
|
||||
&ThreadsafeFunction<String, ()>,
|
||||
&AtomicBool,
|
||||
&AtomicBool,
|
||||
) -> std::result::Result<RoundOutcome, BackendError>,
|
||||
{
|
||||
let mut messages = payload.request.messages.clone();
|
||||
let tool_executor = NapiToolExecutor::new(tool_callback);
|
||||
let event_sink = NapiEventSink::new_with_emitted(callback, emitted);
|
||||
run_tool_loop(
|
||||
&mut messages,
|
||||
max_steps,
|
||||
|messages| {
|
||||
if aborted.load(Ordering::Relaxed) {
|
||||
return Err(backend_transport_error(STREAM_ABORTED_REASON));
|
||||
}
|
||||
|
||||
let request = CoreRequest {
|
||||
messages: messages.to_vec(),
|
||||
stream: true,
|
||||
..payload.request.clone()
|
||||
};
|
||||
|
||||
dispatch_round_fn(&request, callback, &aborted, emitted)
|
||||
},
|
||||
tool_executor,
|
||||
event_sink,
|
||||
|| backend_transport_error("ToolCallLoop max steps reached"),
|
||||
)
|
||||
}
|
||||
|
||||
fn run_native_tool_loop(
|
||||
route: RoutedBackend,
|
||||
payload: LlmDispatchPayload,
|
||||
max_steps: usize,
|
||||
callback: &ThreadsafeFunction<String, ()>,
|
||||
tool_callback: &ThreadsafeFunction<String, PromiseRaw<'static, String>>,
|
||||
aborted: Arc<AtomicBool>,
|
||||
emitted: &AtomicBool,
|
||||
) -> std::result::Result<(), BackendError> {
|
||||
let middleware = payload.middleware.clone();
|
||||
run_native_tool_loop_with_dispatch(
|
||||
payload,
|
||||
max_steps,
|
||||
callback,
|
||||
tool_callback,
|
||||
aborted,
|
||||
emitted,
|
||||
|request, callback, aborted, emitted| dispatch_round(&route, request, callback, &middleware, aborted, emitted),
|
||||
)
|
||||
}
|
||||
|
||||
fn run_native_routed_tool_loop(
|
||||
routes: Vec<RoutedBackend>,
|
||||
payload: LlmDispatchPayload,
|
||||
max_steps: usize,
|
||||
callback: &ThreadsafeFunction<String, ()>,
|
||||
tool_callback: &ThreadsafeFunction<String, PromiseRaw<'static, String>>,
|
||||
aborted: Arc<AtomicBool>,
|
||||
emitted: &AtomicBool,
|
||||
) -> std::result::Result<(), BackendError> {
|
||||
let middleware = payload.middleware.clone();
|
||||
run_native_tool_loop_with_dispatch(
|
||||
payload,
|
||||
max_steps,
|
||||
callback,
|
||||
tool_callback,
|
||||
aborted,
|
||||
emitted,
|
||||
|request, callback, aborted, emitted| {
|
||||
dispatch_round_with_fallback(&routes, request, callback, &middleware, aborted, emitted)
|
||||
},
|
||||
)
|
||||
}
|
||||
|
||||
pub(crate) fn run_native_prepared_tool_loop(
|
||||
routes: Vec<PreparedToolLoopRoute>,
|
||||
max_steps: usize,
|
||||
callback: &ThreadsafeFunction<String, ()>,
|
||||
tool_callback: &ThreadsafeFunction<String, PromiseRaw<'static, String>>,
|
||||
aborted: Arc<AtomicBool>,
|
||||
) -> std::result::Result<(), BackendError> {
|
||||
let Some(((_, request), middleware)) = routes.first() else {
|
||||
return Err(BackendError::NoBackendAvailable);
|
||||
};
|
||||
let payload = LlmDispatchPayload {
|
||||
request: request.clone(),
|
||||
middleware: middleware.clone(),
|
||||
};
|
||||
let emitted = AtomicBool::new(false);
|
||||
|
||||
run_native_tool_loop_with_dispatch(
|
||||
payload,
|
||||
max_steps,
|
||||
callback,
|
||||
tool_callback,
|
||||
aborted,
|
||||
&emitted,
|
||||
|request, callback, aborted, emitted| {
|
||||
dispatch_prepared_payload_round_with_fallback(&routes, request, callback, aborted, emitted)
|
||||
},
|
||||
)
|
||||
}
|
||||
|
||||
pub(crate) fn spawn_tool_loop_stream(
|
||||
protocol: ChatProtocol,
|
||||
config: BackendConfig,
|
||||
payload: LlmDispatchPayload,
|
||||
max_steps: usize,
|
||||
callback: ThreadsafeFunction<String, ()>,
|
||||
tool_callback: ThreadsafeFunction<String, PromiseRaw<'static, String>>,
|
||||
) -> LlmStreamHandle {
|
||||
let aborted = Arc::new(AtomicBool::new(false));
|
||||
let aborted_in_worker = aborted.clone();
|
||||
|
||||
std::thread::spawn(move || {
|
||||
let emitted = AtomicBool::new(false);
|
||||
let result = run_native_tool_loop(
|
||||
RoutedBackend {
|
||||
provider_id: String::new(),
|
||||
protocol,
|
||||
model: payload.request.model.clone(),
|
||||
config,
|
||||
},
|
||||
payload,
|
||||
max_steps,
|
||||
&callback,
|
||||
&tool_callback,
|
||||
aborted_in_worker.clone(),
|
||||
&emitted,
|
||||
);
|
||||
let callback_dispatch_failed = matches!(
|
||||
&result,
|
||||
Err(BackendError::Transport { message: reason })
|
||||
if reason.starts_with(STREAM_CALLBACK_DISPATCH_FAILED_REASON)
|
||||
);
|
||||
|
||||
if let Err(error) = result
|
||||
&& !aborted_in_worker.load(Ordering::Relaxed)
|
||||
&& !matches!(&error, BackendError::Transport { message: reason } if reason == STREAM_ABORTED_REASON)
|
||||
&& !callback_dispatch_failed
|
||||
{
|
||||
emit_error_event(&callback, error.to_string(), "dispatch_error");
|
||||
}
|
||||
|
||||
if !aborted_in_worker.load(Ordering::Relaxed) && !callback_dispatch_failed {
|
||||
let _ = callback.call(
|
||||
Ok(STREAM_END_MARKER.to_string()),
|
||||
ThreadsafeFunctionCallMode::NonBlocking,
|
||||
);
|
||||
}
|
||||
});
|
||||
|
||||
LlmStreamHandle { aborted }
|
||||
}
|
||||
|
||||
pub(crate) fn spawn_routed_tool_loop_stream(
|
||||
routes: Vec<RoutedBackend>,
|
||||
payload: LlmDispatchPayload,
|
||||
max_steps: usize,
|
||||
callback: ThreadsafeFunction<String, ()>,
|
||||
tool_callback: ThreadsafeFunction<String, PromiseRaw<'static, String>>,
|
||||
) -> LlmStreamHandle {
|
||||
let aborted = Arc::new(AtomicBool::new(false));
|
||||
let aborted_in_worker = aborted.clone();
|
||||
|
||||
std::thread::spawn(move || {
|
||||
let emitted = AtomicBool::new(false);
|
||||
let result = run_native_routed_tool_loop(
|
||||
routes,
|
||||
payload,
|
||||
max_steps,
|
||||
&callback,
|
||||
&tool_callback,
|
||||
aborted_in_worker.clone(),
|
||||
&emitted,
|
||||
);
|
||||
let callback_dispatch_failed = matches!(
|
||||
&result,
|
||||
Err(BackendError::Transport { message: reason })
|
||||
if reason.starts_with(STREAM_CALLBACK_DISPATCH_FAILED_REASON)
|
||||
);
|
||||
|
||||
if let Err(error) = result
|
||||
&& !aborted_in_worker.load(Ordering::Relaxed)
|
||||
&& !matches!(&error, BackendError::Transport { message: reason } if reason == STREAM_ABORTED_REASON)
|
||||
&& !callback_dispatch_failed
|
||||
{
|
||||
emit_error_event(&callback, error.to_string(), "dispatch_error");
|
||||
}
|
||||
|
||||
if !aborted_in_worker.load(Ordering::Relaxed) && !callback_dispatch_failed {
|
||||
let _ = callback.call(
|
||||
Ok(STREAM_END_MARKER.to_string()),
|
||||
ThreadsafeFunctionCallMode::NonBlocking,
|
||||
);
|
||||
}
|
||||
});
|
||||
|
||||
LlmStreamHandle { aborted }
|
||||
}
|
||||
|
||||
pub(crate) fn spawn_prepared_tool_loop_stream(
|
||||
routes: Vec<PreparedToolLoopRoute>,
|
||||
max_steps: usize,
|
||||
callback: ThreadsafeFunction<String, ()>,
|
||||
tool_callback: ThreadsafeFunction<String, PromiseRaw<'static, String>>,
|
||||
) -> LlmStreamHandle {
|
||||
let aborted = Arc::new(AtomicBool::new(false));
|
||||
let aborted_in_worker = aborted.clone();
|
||||
|
||||
std::thread::spawn(move || {
|
||||
let result = run_native_prepared_tool_loop(routes, max_steps, &callback, &tool_callback, aborted_in_worker.clone());
|
||||
let callback_dispatch_failed = matches!(
|
||||
&result,
|
||||
Err(BackendError::Transport { message: reason })
|
||||
if reason.starts_with(STREAM_CALLBACK_DISPATCH_FAILED_REASON)
|
||||
);
|
||||
|
||||
if let Err(error) = result
|
||||
&& !aborted_in_worker.load(Ordering::Relaxed)
|
||||
&& !matches!(&error, BackendError::Transport { message: reason } if reason == STREAM_ABORTED_REASON)
|
||||
&& !callback_dispatch_failed
|
||||
{
|
||||
emit_error_event(&callback, error.to_string(), "dispatch_error");
|
||||
}
|
||||
|
||||
if !aborted_in_worker.load(Ordering::Relaxed) && !callback_dispatch_failed {
|
||||
let _ = callback.call(
|
||||
Ok(STREAM_END_MARKER.to_string()),
|
||||
ThreadsafeFunctionCallMode::NonBlocking,
|
||||
);
|
||||
}
|
||||
});
|
||||
|
||||
LlmStreamHandle { aborted }
|
||||
}
|
||||
@@ -0,0 +1,8 @@
|
||||
mod callback;
|
||||
mod contract;
|
||||
mod engine;
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests;
|
||||
|
||||
pub(crate) use engine::{spawn_prepared_tool_loop_stream, spawn_routed_tool_loop_stream, spawn_tool_loop_stream};
|
||||
@@ -0,0 +1,36 @@
|
||||
use llm_adapter::core::{CoreContent, CoreMessage};
|
||||
use llm_runtime::{ToolResultMessage, append_tool_turns};
|
||||
use serde_json::json;
|
||||
|
||||
use super::contract::NativeToolCall;
|
||||
|
||||
#[test]
|
||||
fn append_tool_turns_should_replay_assistant_and_tool_messages() {
|
||||
let mut messages = vec![CoreMessage {
|
||||
role: llm_adapter::core::CoreRole::User,
|
||||
content: vec![CoreContent::Text {
|
||||
text: "read doc".to_string(),
|
||||
}],
|
||||
}];
|
||||
|
||||
append_tool_turns(
|
||||
&mut messages,
|
||||
&[NativeToolCall {
|
||||
id: "call_1".to_string(),
|
||||
name: "doc_read".to_string(),
|
||||
args: json!({ "doc_id": "a1" }),
|
||||
raw_arguments_text: Some("{\"doc_id\":\"a1\"}".to_string()),
|
||||
argument_parse_error: None,
|
||||
thought: Some("need context".to_string()),
|
||||
}],
|
||||
&[ToolResultMessage {
|
||||
call_id: "call_1".to_string(),
|
||||
output: json!({ "markdown": "# doc" }),
|
||||
is_error: Some(false),
|
||||
}],
|
||||
);
|
||||
|
||||
assert_eq!(messages.len(), 3);
|
||||
assert!(matches!(messages[1].role, llm_adapter::core::CoreRole::Assistant));
|
||||
assert!(matches!(messages[2].role, llm_adapter::core::CoreRole::Tool));
|
||||
}
|
||||
Reference in New Issue
Block a user