mirror of
https://github.com/toeverything/AFFiNE.git
synced 2026-09-02 06:39:46 +08:00
feat(server): refactor for byok (#14911)
This commit is contained in:
@@ -7,7 +7,7 @@ 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(crate) use stream::{emit_error_event, emit_provider_selected_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,
|
||||
|
||||
@@ -106,14 +106,18 @@ fn spawn_prepared_stream(
|
||||
if reason.starts_with(STREAM_CALLBACK_DISPATCH_FAILED_REASON)
|
||||
);
|
||||
|
||||
if let Err(error) = result
|
||||
if let Err(error) = &result
|
||||
&& !aborted_in_worker.load(Ordering::Relaxed)
|
||||
&& !callback_dispatch_failed
|
||||
&& !is_abort_error(&error)
|
||||
&& !is_abort_error(error)
|
||||
{
|
||||
emit_error_event(&callback, error.to_string(), "dispatch_error");
|
||||
}
|
||||
|
||||
if let Ok(provider_id) = result {
|
||||
emit_provider_selected_event(&callback, provider_id);
|
||||
}
|
||||
|
||||
if !callback_dispatch_failed {
|
||||
let _ = callback.call(
|
||||
Ok(STREAM_END_MARKER.to_string()),
|
||||
@@ -129,7 +133,7 @@ fn dispatch_prepared_stream_with_fallback(
|
||||
routes: &[PreparedDispatchRoute],
|
||||
callback: &ThreadsafeFunction<String, ()>,
|
||||
aborted: &AtomicBool,
|
||||
) -> std::result::Result<(), BackendError> {
|
||||
) -> std::result::Result<String, BackendError> {
|
||||
dispatch_prepared_stream_with_fallback_using_client(&DefaultHttpClient::default(), routes, aborted, |event| {
|
||||
emit_stream_event(callback, event)
|
||||
})
|
||||
@@ -140,7 +144,7 @@ fn dispatch_prepared_stream_with_fallback_using_client<F>(
|
||||
routes: &[PreparedDispatchRoute],
|
||||
aborted: &AtomicBool,
|
||||
mut emit_event: F,
|
||||
) -> std::result::Result<(), BackendError>
|
||||
) -> std::result::Result<String, BackendError>
|
||||
where
|
||||
F: FnMut(&StreamEvent) -> Status,
|
||||
{
|
||||
@@ -154,7 +158,7 @@ where
|
||||
.collect::<std::result::Result<Vec<_>, BackendError>>()?;
|
||||
let mut callback_dispatch_failed = false;
|
||||
|
||||
dispatch_prepared_stream_with_pipeline(
|
||||
let provider_id = dispatch_prepared_stream_with_pipeline(
|
||||
client,
|
||||
&mut adapter_routes,
|
||||
|| aborted.load(Ordering::Relaxed),
|
||||
@@ -174,7 +178,7 @@ where
|
||||
"{STREAM_CALLBACK_DISPATCH_FAILED_REASON}:unknown"
|
||||
)))
|
||||
} else {
|
||||
Ok(())
|
||||
Ok(provider_id)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -195,6 +199,16 @@ pub(crate) fn emit_error_event(callback: &ThreadsafeFunction<String, ()>, messag
|
||||
let _ = callback.call(Ok(error_event), ThreadsafeFunctionCallMode::NonBlocking);
|
||||
}
|
||||
|
||||
pub(crate) fn emit_provider_selected_event(callback: &ThreadsafeFunction<String, ()>, provider_id: String) {
|
||||
let event = serde_json::json!({
|
||||
"type": "provider_selected",
|
||||
"provider_id": provider_id,
|
||||
})
|
||||
.to_string();
|
||||
|
||||
let _ = callback.call(Ok(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!({
|
||||
|
||||
@@ -14,7 +14,10 @@ use napi::{
|
||||
threadsafe_function::{ThreadsafeFunction, ThreadsafeFunctionCallMode},
|
||||
};
|
||||
|
||||
use super::callback::{NapiEventSink, NapiToolExecutor, emit_tool_loop_event};
|
||||
use super::{
|
||||
super::emit_provider_selected_event,
|
||||
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,
|
||||
@@ -39,11 +42,13 @@ fn dispatch_prepared_round_with_fallback(
|
||||
})
|
||||
.collect::<std::result::Result<Vec<_>, BackendError>>()?;
|
||||
|
||||
run_prepared_stream_round_with_fallback(
|
||||
let mut selected_provider_id: Option<String> = None;
|
||||
let outcome = run_prepared_stream_round_with_fallback(
|
||||
&mut pipelines,
|
||||
|on_event| {
|
||||
let (selected_index, _) =
|
||||
let (selected_index, provider_id) =
|
||||
dispatch_prepared_stream_with_fallback_index(&DefaultHttpClient::default(), &adapter_routes, on_event)?;
|
||||
selected_provider_id = Some(provider_id);
|
||||
Ok(selected_index)
|
||||
},
|
||||
|| aborted.load(Ordering::Relaxed),
|
||||
@@ -53,7 +58,11 @@ fn dispatch_prepared_round_with_fallback(
|
||||
emitted.store(true, Ordering::Relaxed);
|
||||
emit_tool_loop_event(callback, loop_event)
|
||||
},
|
||||
)
|
||||
)?;
|
||||
if let Some(provider_id) = selected_provider_id {
|
||||
emit_provider_selected_event(callback, provider_id);
|
||||
}
|
||||
Ok(outcome)
|
||||
}
|
||||
|
||||
fn prepare_tool_loop_route(
|
||||
|
||||
Reference in New Issue
Block a user