first pass of merging in warp (doesn't build)
This commit is contained in:
@@ -0,0 +1,170 @@
|
||||
use base64::Engine as _;
|
||||
use base64::prelude::BASE64_URL_SAFE;
|
||||
use futures::StreamExt as _;
|
||||
use prost::Message as _;
|
||||
use tracing_futures::Instrument as _;
|
||||
use warp_core::channel::ChannelState;
|
||||
#[cfg(feature = "agent_mode_evals")]
|
||||
use warp_server_client::base_client::EVAL_USER_ID_HEADER;
|
||||
use warp_server_client::base_client::{AmbientHeaderPolicy, BaseClient};
|
||||
|
||||
#[derive(Debug, thiserror::Error)]
|
||||
pub enum Error {
|
||||
#[error("Failed to authenticate multi-agent request")]
|
||||
Authentication(#[source] anyhow::Error),
|
||||
|
||||
#[error("Failed to resolve ambient headers for multi-agent request")]
|
||||
AmbientHeaders(#[source] anyhow::Error),
|
||||
|
||||
#[error("Failed to decode base64 multi-agent response event")]
|
||||
Base64Decode(#[source] base64::DecodeError),
|
||||
|
||||
#[error("Failed to decode protobuf multi-agent response event")]
|
||||
ProtobufDecode(#[source] prost::DecodeError),
|
||||
|
||||
#[error("Multi-agent eventsource stream failed: {0:?}")]
|
||||
EventSource(Box<reqwest_eventsource::Error>),
|
||||
}
|
||||
|
||||
cfg_if::cfg_if! {
|
||||
if #[cfg(target_family = "wasm")] {
|
||||
/// A multi-agent response event stream without an unnecessary `Send` bound on WASM.
|
||||
pub type OutputStream = futures::stream::LocalBoxStream<
|
||||
'static,
|
||||
Result<warp_multi_agent_api::ResponseEvent, Error>,
|
||||
>;
|
||||
} else {
|
||||
/// A multi-agent response event stream that can be sent between native threads.
|
||||
pub type OutputStream = futures::stream::BoxStream<
|
||||
'static,
|
||||
Result<warp_multi_agent_api::ResponseEvent, Error>,
|
||||
>;
|
||||
}
|
||||
}
|
||||
|
||||
/// Opens a decoded multi-agent response event stream.
|
||||
pub async fn generate_multi_agent_output(
|
||||
client: &BaseClient,
|
||||
request: &warp_multi_agent_api::Request,
|
||||
) -> Result<OutputStream, Error> {
|
||||
let auth_token = client
|
||||
.get_or_refresh_access_token()
|
||||
.await
|
||||
.map_err(Error::Authentication)?;
|
||||
let is_passive = is_passive_suggestion_request(request);
|
||||
let url = endpoint_url(is_passive);
|
||||
|
||||
let mut request_builder = client
|
||||
.http_client()
|
||||
.post(url)
|
||||
.proto(request)
|
||||
.prevent_sleep("Agent Mode request in-progress");
|
||||
if let Some(token) = auth_token.as_bearer_token() {
|
||||
request_builder = request_builder.bearer_auth(token);
|
||||
}
|
||||
|
||||
for (name, value) in client
|
||||
.ambient_headers(ambient_policy(is_passive))
|
||||
.await
|
||||
.map_err(Error::AmbientHeaders)?
|
||||
{
|
||||
request_builder = request_builder.header(name, value);
|
||||
}
|
||||
|
||||
#[cfg(feature = "agent_mode_evals")]
|
||||
if let Some(eval_user_id) = client.eval_user_id() {
|
||||
request_builder = request_builder.header(EVAL_USER_ID_HEADER, eval_user_id.to_string());
|
||||
}
|
||||
|
||||
let raw_stream = client.wrap_eventsource_with_iap_detection(request_builder.eventsource());
|
||||
let output_stream = raw_stream.filter_map(|event| async {
|
||||
match event {
|
||||
Ok(reqwest_eventsource::Event::Message(message_event)) => {
|
||||
Some(decode_response_event(&message_event.data))
|
||||
}
|
||||
Ok(reqwest_eventsource::Event::Open) => None,
|
||||
Err(error) => Some(Err(Error::EventSource(Box::new(error)))),
|
||||
}
|
||||
});
|
||||
|
||||
// Once we get the init event, add some identifiers to the trace span.
|
||||
let output_stream = output_stream.inspect(|event| {
|
||||
if let Ok(event) = &event {
|
||||
match &event.r#type {
|
||||
Some(warp_multi_agent_api::response_event::Type::Init(init)) => {
|
||||
tracing::info!("StreamInit");
|
||||
tracing::Span::current().record("conversation_id", &init.conversation_id);
|
||||
tracing::Span::current().record("request_id", &init.request_id);
|
||||
tracing::Span::current().record("run_id", &init.run_id);
|
||||
}
|
||||
Some(warp_multi_agent_api::response_event::Type::Finished(_finished)) => {
|
||||
tracing::info!("StreamFinished");
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
});
|
||||
// Wrap the output stream with a trace span.
|
||||
let output_stream = output_stream.instrument(tracing::info_span!(
|
||||
"generate_multi_agent_output",
|
||||
tags.cloud_agent = true,
|
||||
conversation_id = tracing::field::Empty,
|
||||
request_id = tracing::field::Empty,
|
||||
run_id = tracing::field::Empty,
|
||||
));
|
||||
|
||||
cfg_if::cfg_if! {
|
||||
if #[cfg(target_family = "wasm")] {
|
||||
Ok(output_stream.boxed_local())
|
||||
} else {
|
||||
Ok(output_stream.boxed())
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn is_passive_suggestion_request(request: &warp_multi_agent_api::Request) -> bool {
|
||||
request.input.as_ref().is_some_and(|input| {
|
||||
matches!(
|
||||
input.r#type,
|
||||
Some(warp_multi_agent_api::request::input::Type::GeneratePassiveSuggestions(_))
|
||||
)
|
||||
})
|
||||
}
|
||||
|
||||
fn endpoint_url(is_passive: bool) -> String {
|
||||
format!(
|
||||
"{}/{}/{}",
|
||||
ChannelState::server_root_url(),
|
||||
if cfg!(feature = "agent_mode_evals") {
|
||||
"agent-mode-evals"
|
||||
} else {
|
||||
"ai"
|
||||
},
|
||||
if is_passive {
|
||||
"passive-suggestions"
|
||||
} else {
|
||||
"multi-agent"
|
||||
}
|
||||
)
|
||||
}
|
||||
|
||||
fn ambient_policy(is_passive: bool) -> AmbientHeaderPolicy {
|
||||
if is_passive {
|
||||
// Passive suggestions read from the main conversation, but cannot modify it.
|
||||
AmbientHeaderPolicy::omit_all()
|
||||
} else {
|
||||
AmbientHeaderPolicy::workload_only()
|
||||
}
|
||||
}
|
||||
|
||||
fn decode_response_event(data: &str) -> Result<warp_multi_agent_api::ResponseEvent, Error> {
|
||||
let decoded_data = BASE64_URL_SAFE
|
||||
.decode(data.trim_matches('"'))
|
||||
.map_err(Error::Base64Decode)?;
|
||||
warp_multi_agent_api::ResponseEvent::decode(decoded_data.as_slice())
|
||||
.map_err(Error::ProtobufDecode)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
#[path = "lib_tests.rs"]
|
||||
mod tests;
|
||||
@@ -0,0 +1,77 @@
|
||||
use base64::Engine as _;
|
||||
use base64::prelude::BASE64_URL_SAFE;
|
||||
use prost::Message as _;
|
||||
use warp_server_client::base_client::AmbientHeaderPolicy;
|
||||
|
||||
use super::{
|
||||
Error, ambient_policy, decode_response_event, endpoint_url, is_passive_suggestion_request,
|
||||
};
|
||||
|
||||
#[test]
|
||||
fn detects_passive_suggestion_requests() {
|
||||
let regular = warp_multi_agent_api::Request::default();
|
||||
let passive = warp_multi_agent_api::Request {
|
||||
input: Some(warp_multi_agent_api::request::Input {
|
||||
r#type: Some(
|
||||
warp_multi_agent_api::request::input::Type::GeneratePassiveSuggestions(
|
||||
Default::default(),
|
||||
),
|
||||
),
|
||||
..Default::default()
|
||||
}),
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
assert!(!is_passive_suggestion_request(®ular));
|
||||
assert!(is_passive_suggestion_request(&passive));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn routes_regular_and_passive_requests_to_distinct_endpoints() {
|
||||
let prefix = if cfg!(feature = "agent_mode_evals") {
|
||||
"agent-mode-evals"
|
||||
} else {
|
||||
"ai"
|
||||
};
|
||||
|
||||
assert!(endpoint_url(false).ends_with(&format!("/{prefix}/multi-agent")));
|
||||
assert!(endpoint_url(true).ends_with(&format!("/{prefix}/passive-suggestions")));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn selects_endpoint_specific_ambient_header_policies() {
|
||||
assert_eq!(ambient_policy(false), AmbientHeaderPolicy::workload_only());
|
||||
assert_eq!(ambient_policy(true), AmbientHeaderPolicy::omit_all());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn decodes_quoted_base64_protobuf_response_event() {
|
||||
let expected = warp_multi_agent_api::ResponseEvent::default();
|
||||
let encoded = BASE64_URL_SAFE.encode(expected.encode_to_vec());
|
||||
|
||||
let decoded = decode_response_event(&format!("\"{encoded}\"")).unwrap();
|
||||
|
||||
assert_eq!(decoded, expected);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn distinguishes_base64_and_protobuf_decode_errors() {
|
||||
assert!(matches!(
|
||||
decode_response_event("%"),
|
||||
Err(Error::Base64Decode(_))
|
||||
));
|
||||
|
||||
let invalid_protobuf = BASE64_URL_SAFE.encode([0xff]);
|
||||
assert!(matches!(
|
||||
decode_response_event(&invalid_protobuf),
|
||||
Err(Error::ProtobufDecode(_))
|
||||
));
|
||||
}
|
||||
|
||||
#[cfg(not(target_family = "wasm"))]
|
||||
#[test]
|
||||
fn native_output_stream_is_send() {
|
||||
fn assert_send<T: Send>() {}
|
||||
|
||||
assert_send::<super::OutputStream>();
|
||||
}
|
||||
Reference in New Issue
Block a user