add support for json based content types in chat completion request

This commit is contained in:
Adil Hafeez 2025-05-22 16:11:04 -07:00
parent 27c0f2fdce
commit d089ef0ed8
No known key found for this signature in database
GPG key ID: 9B18EF7691369645
14 changed files with 309 additions and 111 deletions

View file

@ -625,7 +625,7 @@ static_resources:
- endpoint: - endpoint:
address: address:
socket_address: socket_address:
address: 0.0.0.0 address: host.docker.internal
port_value: 9091 port_value: 9091
hostname: localhost hostname: localhost

View file

@ -52,6 +52,8 @@ def docker_start_archgw_detached(
port_mappings = [ port_mappings = [
f"{prompt_gateway_port}:{prompt_gateway_port}", f"{prompt_gateway_port}:{prompt_gateway_port}",
f"{llm_gateway_port}:{llm_gateway_port}", f"{llm_gateway_port}:{llm_gateway_port}",
# expose llm gateway as well without tracing support
f"{llm_gateway_port + 1}:{llm_gateway_port + 1}",
"9901:19901", "9901:19901",
] ]
port_mappings_args = [item for port in port_mappings for item in ("-p", port)] port_mappings_args = [item for port in port_mappings for item in ("-p", port)]

View file

@ -9,10 +9,11 @@ use http_body_util::{BodyExt, Full, StreamBody};
use hyper::body::Frame; use hyper::body::Frame;
use hyper::header::{self}; use hyper::header::{self};
use hyper::{Request, Response, StatusCode}; use hyper::{Request, Response, StatusCode};
use serde_json::Value;
use tokio::sync::mpsc; use tokio::sync::mpsc;
use tokio_stream::wrappers::ReceiverStream; use tokio_stream::wrappers::ReceiverStream;
use tokio_stream::StreamExt; use tokio_stream::StreamExt;
use tracing::{info, warn}; use tracing::{debug, info, warn};
use crate::router::llm_router::RouterService; use crate::router::llm_router::RouterService;
@ -30,19 +31,23 @@ pub async fn chat_completions(
let mut request_headers = request.headers().clone(); let mut request_headers = request.headers().clone();
let chat_request_bytes = request.collect().await?.to_bytes(); let chat_request_bytes = request.collect().await?.to_bytes();
let chat_completion_request: ChatCompletionsRequest = let chat_completion_request: ChatCompletionsRequest =
match serde_json::from_slice(&chat_request_bytes) { match serde_json::from_slice(&chat_request_bytes) {
Ok(request) => request, Ok(request) => request,
Err(err) => { Err(err) => {
let v: Value = serde_json::from_slice(&chat_request_bytes).unwrap();
let err_msg = format!("Failed to parse request body: {}", err); let err_msg = format!("Failed to parse request body: {}", err);
warn!("{}", err_msg);
warn!("request body: {}", v.to_string());
let mut bad_request = Response::new(full(err_msg)); let mut bad_request = Response::new(full(err_msg));
*bad_request.status_mut() = StatusCode::BAD_REQUEST; *bad_request.status_mut() = StatusCode::BAD_REQUEST;
return Ok(bad_request); return Ok(bad_request);
} }
}; };
info!( debug!(
"request body received: {}", "request body: {}",
shorten_string(&serde_json::to_string(&chat_completion_request).unwrap()) shorten_string(&serde_json::to_string(&chat_completion_request).unwrap())
); );

View file

@ -2,38 +2,27 @@ use brightstaff::handlers::chat_completions::chat_completions;
use brightstaff::router::llm_router::RouterService; use brightstaff::router::llm_router::RouterService;
use bytes::Bytes; use bytes::Bytes;
use common::configuration::Configuration; use common::configuration::Configuration;
use common::utils::shorten_string;
use http_body_util::{combinators::BoxBody, BodyExt, Empty}; use http_body_util::{combinators::BoxBody, BodyExt, Empty};
use hyper::body::Incoming; use hyper::body::Incoming;
use hyper::server::conn::http1; use hyper::server::conn::http1;
use hyper::service::service_fn; use hyper::service::service_fn;
use hyper::{Method, Request, Response, StatusCode}; use hyper::{Method, Request, Response, StatusCode};
use hyper_util::rt::TokioIo; use hyper_util::rt::TokioIo;
use opentelemetry::global::BoxedTracer;
use opentelemetry::trace::FutureExt; use opentelemetry::trace::FutureExt;
use opentelemetry::{ use opentelemetry::{global, Context};
global,
trace::{SpanKind, Tracer},
Context,
};
use opentelemetry_http::HeaderExtractor; use opentelemetry_http::HeaderExtractor;
use opentelemetry_sdk::{propagation::TraceContextPropagator, trace::SdkTracerProvider}; use opentelemetry_sdk::{propagation::TraceContextPropagator, trace::SdkTracerProvider};
use opentelemetry_stdout::SpanExporter; use opentelemetry_stdout::SpanExporter;
use std::sync::{Arc, OnceLock}; use std::sync::Arc;
use std::{env, fs}; use std::{env, fs};
use tokio::net::TcpListener; use tokio::net::TcpListener;
use tracing::info; use tracing::{debug, info};
use tracing_subscriber::EnvFilter; use tracing_subscriber::EnvFilter;
pub mod router; pub mod router;
const BIND_ADDRESS: &str = "0.0.0.0:9091"; const BIND_ADDRESS: &str = "0.0.0.0:9091";
fn get_tracer() -> &'static BoxedTracer {
static TRACER: OnceLock<BoxedTracer> = OnceLock::new();
TRACER.get_or_init(|| global::tracer("archgw/router"))
}
// Utility function to extract the context from the incoming request headers // Utility function to extract the context from the incoming request headers
fn extract_context_from_request(req: &Request<Incoming>) -> Context { fn extract_context_from_request(req: &Request<Incoming>) -> Context {
global::get_text_map_propagator(|propagator| { global::get_text_map_propagator(|propagator| {
@ -83,24 +72,23 @@ async fn main() -> Result<(), Box<dyn std::error::Error + Send + Sync>> {
let arch_config = Arc::new(config); let arch_config = Arc::new(config);
info!( debug!(
"arch_config: {:?}", "arch_config: {:?}",
shorten_string(&serde_json::to_string(arch_config.as_ref()).unwrap()) &serde_json::to_string(arch_config.as_ref()).unwrap()
); );
let llm_provider_endpoint = env::var("LLM_PROVIDER_ENDPOINT") let llm_provider_endpoint = env::var("LLM_PROVIDER_ENDPOINT")
.unwrap_or_else(|_| "http://localhost:12001/v1/chat/completions".to_string()); .unwrap_or_else(|_| "http://localhost:12001/v1/chat/completions".to_string());
info!("llm provider endpoint: {}", llm_provider_endpoint); info!("llm provider endpoint: {}", llm_provider_endpoint);
info!("Listening on http://{}", bind_address); info!("listening on http://{}", bind_address);
let listener = TcpListener::bind(bind_address).await?; let listener = TcpListener::bind(bind_address).await?;
// if routing is null then return gpt-4o as model name // if routing is null then return gpt-4o as model name
let model = arch_config.routing.as_ref().map_or_else( let model = arch_config
|| "gpt-4o".to_string(), .routing
|routing| routing.model.clone(), .as_ref()
); .map_or_else(|| "gpt-4o".to_string(), |routing| routing.model.clone());
let router_service: Arc<RouterService> = Arc::new(RouterService::new( let router_service: Arc<RouterService> = Arc::new(RouterService::new(
arch_config.llm_providers.clone(), arch_config.llm_providers.clone(),
@ -119,12 +107,6 @@ async fn main() -> Result<(), Box<dyn std::error::Error + Send + Sync>> {
let service = service_fn(move |req| { let service = service_fn(move |req| {
let router_service = Arc::clone(&router_service); let router_service = Arc::clone(&router_service);
let parent_cx = extract_context_from_request(&req); let parent_cx = extract_context_from_request(&req);
info!("parent_cx: {:?}", parent_cx);
let tracer = get_tracer();
let _span = tracer
.span_builder("request")
.with_kind(SpanKind::Server)
.start_with_context(tracer, &parent_cx);
let llm_provider_endpoint = llm_provider_endpoint.clone(); let llm_provider_endpoint = llm_provider_endpoint.clone();
async move { async move {

View file

@ -1,14 +1,13 @@
use std::sync::Arc; use std::sync::Arc;
use common::{ use common::{
api::open_ai::{ChatCompletionsResponse, Message}, api::open_ai::{ChatCompletionsResponse, ContentType, Message},
configuration::LlmProvider, configuration::LlmProvider,
consts::ARCH_PROVIDER_HINT_HEADER, consts::ARCH_PROVIDER_HINT_HEADER,
utils::shorten_string,
}; };
use hyper::header; use hyper::header;
use thiserror::Error; use thiserror::Error;
use tracing::{info, warn}; use tracing::{debug, info, warn};
use super::router_model::RouterModel; use super::router_model::RouterModel;
@ -59,9 +58,9 @@ impl RouterService {
.collect::<Vec<String>>() .collect::<Vec<String>>()
.join("\n"); .join("\n");
info!( debug!(
"llm_providers from config with usage: {}...", "llm_providers from config with usage: {}...",
shorten_string(&llm_providers_with_usage_yaml.replace("\n", "\\n")) llm_providers_with_usage_yaml.replace("\n", "\\n")
); );
let router_model = Arc::new(super::router_model_v1::RouterModelV1::new( let router_model = Arc::new(super::router_model_v1::RouterModelV1::new(
@ -83,7 +82,6 @@ impl RouterService {
messages: &[Message], messages: &[Message],
trace_parent: Option<String>, trace_parent: Option<String>,
) -> Result<Option<String>> { ) -> Result<Option<String>> {
if !self.llm_usage_defined { if !self.llm_usage_defined {
return Ok(None); return Ok(None);
} }
@ -91,8 +89,14 @@ impl RouterService {
let router_request = self.router_model.generate_request(messages); let router_request = self.router_model.generate_request(messages);
info!( info!(
"router_request: {}", "sending request to arch-router model: {}, endpoint: {}",
shorten_string(&serde_json::to_string(&router_request).unwrap()), self.router_model.get_model_name(),
self.router_url
);
debug!(
"arch request body: {}",
&serde_json::to_string(&router_request).unwrap(),
); );
let mut llm_route_request_headers = header::HeaderMap::new(); let mut llm_route_request_headers = header::HeaderMap::new();
@ -113,6 +117,7 @@ impl RouterService {
); );
} }
let start_time = std::time::Instant::now();
let res = self let res = self
.client .client
.post(&self.router_url) .post(&self.router_url)
@ -122,6 +127,7 @@ impl RouterService {
.await?; .await?;
let body = res.text().await?; let body = res.text().await?;
let router_response_time = start_time.elapsed();
let chat_completion_response: ChatCompletionsResponse = match serde_json::from_str(&body) { let chat_completion_response: ChatCompletionsResponse = match serde_json::from_str(&body) {
Ok(response) => response, Ok(response) => response,
@ -138,14 +144,18 @@ impl RouterService {
} }
}; };
let selected_llm = self.router_model.parse_response( if let Some(ContentType::Text(content)) =
chat_completion_response.choices[0] &chat_completion_response.choices[0].message.content
.message {
.content info!(
.as_ref() "router response: {}, response time: {}ms",
.unwrap(), content.replace("\n", "\\n"),
)?; router_response_time.as_millis()
);
Ok(selected_llm) let selected_llm = self.router_model.parse_response(content)?;
Ok(selected_llm)
} else {
Ok(None)
}
} }
} }

View file

@ -12,4 +12,5 @@ pub type Result<T> = std::result::Result<T, RoutingModelError>;
pub trait RouterModel: Send + Sync { pub trait RouterModel: Send + Sync {
fn generate_request(&self, messages: &[Message]) -> ChatCompletionsRequest; fn generate_request(&self, messages: &[Message]) -> ChatCompletionsRequest;
fn parse_response(&self, content: &str) -> Result<Option<String>>; fn parse_response(&self, content: &str) -> Result<Option<String>>;
fn get_model_name(&self) -> String;
} }

View file

@ -1,9 +1,8 @@
use common::{ use common::{
api::open_ai::{ChatCompletionsRequest, Message}, api::open_ai::{ChatCompletionsRequest, ContentType, Message},
consts::{SYSTEM_ROLE, USER_ROLE}, consts::{SYSTEM_ROLE, USER_ROLE},
}; };
use serde::{Deserialize, Serialize}; use serde::{Deserialize, Serialize};
use tracing::info;
use super::router_model::{RouterModel, RoutingModelError}; use super::router_model::{RouterModel, RoutingModelError};
@ -68,7 +67,7 @@ impl RouterModel for RouterModelV1 {
ChatCompletionsRequest { ChatCompletionsRequest {
model: self.routing_model.clone(), model: self.routing_model.clone(),
messages: vec![Message { messages: vec![Message {
content: Some(message), content: Some(ContentType::Text(message)),
role: USER_ROLE.to_string(), role: USER_ROLE.to_string(),
model: None, model: None,
tool_calls: None, tool_calls: None,
@ -86,10 +85,6 @@ impl RouterModel for RouterModelV1 {
return Ok(None); return Ok(None);
} }
let router_resp_fixed = fix_json_response(content); let router_resp_fixed = fix_json_response(content);
info!(
"router response (fixed): {}",
router_resp_fixed.replace("\n", "\\n")
);
let router_response: LlmRouterResponse = serde_json::from_str(router_resp_fixed.as_str())?; let router_response: LlmRouterResponse = serde_json::from_str(router_resp_fixed.as_str())?;
let selected_llm = router_response.route.unwrap_or_default().to_string(); let selected_llm = router_response.route.unwrap_or_default().to_string();
@ -100,6 +95,10 @@ impl RouterModel for RouterModelV1 {
Ok(Some(selected_llm)) Ok(Some(selected_llm))
} }
fn get_model_name(&self) -> String {
self.routing_model.clone()
}
} }
fn fix_json_response(body: &str) -> String { fn fix_json_response(body: &str) -> String {
@ -172,22 +171,28 @@ user: "seattle"
let messages = vec![ let messages = vec![
Message { Message {
role: "system".to_string(), role: "system".to_string(),
content: Some("You are a helpful assistant.".to_string()), content: Some(ContentType::Text(
"You are a helpful assistant.".to_string(),
)),
..Default::default() ..Default::default()
}, },
Message { Message {
role: "user".to_string(), role: "user".to_string(),
content: Some("Hello, I want to book a flight.".to_string()), content: Some(ContentType::Text(
"Hello, I want to book a flight.".to_string(),
)),
..Default::default() ..Default::default()
}, },
Message { Message {
role: "assistant".to_string(), role: "assistant".to_string(),
content: Some("Sure, where would you like to go?".to_string()), content: Some(ContentType::Text(
"Sure, where would you like to go?".to_string(),
)),
..Default::default() ..Default::default()
}, },
Message { Message {
role: "user".to_string(), role: "user".to_string(),
content: Some("seattle".to_string()), content: Some(ContentType::Text("seattle".to_string())),
..Default::default() ..Default::default()
}, },
]; ];
@ -198,7 +203,7 @@ user: "seattle"
println!("Prompt: {}", prompt); println!("Prompt: {}", prompt);
assert_eq!(expected_prompt, prompt); assert_eq!(expected_prompt, prompt.to_string());
} }
} }

View file

@ -6,6 +6,8 @@ use crate::{
}; };
use serde::{Deserialize, Serialize}; use serde::{Deserialize, Serialize};
use super::open_ai::ContentType;
#[derive(Debug, Clone, Serialize, Deserialize)] #[derive(Debug, Clone, Serialize, Deserialize)]
pub struct HallucinationClassificationRequest { pub struct HallucinationClassificationRequest {
pub prompt: String, pub prompt: String,
@ -21,7 +23,7 @@ pub struct HallucinationClassificationResponse {
pub fn extract_messages_for_hallucination(messages: &[Message]) -> Vec<String> { pub fn extract_messages_for_hallucination(messages: &[Message]) -> Vec<String> {
let mut arch_assistant = false; let mut arch_assistant = false;
let mut user_messages = Vec::new(); let mut user_messages: Vec<String> = Vec::new();
if messages.len() >= 2 { if messages.len() >= 2 {
let latest_assistant_message = &messages[messages.len() - 2]; let latest_assistant_message = &messages[messages.len() - 2];
if let Some(model) = latest_assistant_message.model.as_ref() { if let Some(model) = latest_assistant_message.model.as_ref() {
@ -34,7 +36,7 @@ pub fn extract_messages_for_hallucination(messages: &[Message]) -> Vec<String> {
for message in messages.iter().rev() { for message in messages.iter().rev() {
if let Some(model) = message.model.as_ref() { if let Some(model) = message.model.as_ref() {
if !model.starts_with(ARCH_MODEL_PREFIX) { if !model.starts_with(ARCH_MODEL_PREFIX) {
if let Some(content) = &message.content { if let Some(ContentType::Text(content)) = &message.content {
if !content.starts_with(HALLUCINATION_TEMPLATE) { if !content.starts_with(HALLUCINATION_TEMPLATE) {
break; break;
} }
@ -42,13 +44,13 @@ pub fn extract_messages_for_hallucination(messages: &[Message]) -> Vec<String> {
} }
} }
if message.role == USER_ROLE { if message.role == USER_ROLE {
if let Some(content) = &message.content { if let Some(ContentType::Text(content)) = &message.content {
user_messages.push(content.clone()); user_messages.push(content.clone());
} }
} }
} }
} else if let Some(message) = messages.last() { } else if let Some(message) = messages.last() {
if let Some(content) = &message.content { if let Some(ContentType::Text(content)) = &message.content {
user_messages.push(content.clone()); user_messages.push(content.clone());
} }
} }

View file

@ -1,6 +1,7 @@
use crate::consts::{ARCH_FC_MODEL_NAME, ASSISTANT_ROLE}; use crate::consts::{ARCH_FC_MODEL_NAME, ASSISTANT_ROLE};
use serde::{ser::SerializeMap, Deserialize, Serialize}; use serde::{ser::SerializeMap, Deserialize, Serialize};
use serde_yaml::Value; use serde_yaml::Value;
use core::panic;
use std::{ use std::{
collections::{HashMap, VecDeque}, collections::{HashMap, VecDeque},
fmt::Display, fmt::Display,
@ -154,12 +155,54 @@ pub struct StreamOptions {
pub include_usage: bool, pub include_usage: bool,
} }
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
pub enum MultiPartContentType {
#[serde(rename = "text")]
Text,
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
pub struct MultiPartContent {
pub text: Option<String>,
#[serde(rename = "type")]
pub content_type: MultiPartContentType,
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
#[serde(untagged)]
pub enum ContentType {
Text(String),
MultiPart(Vec<MultiPartContent>),
}
impl Display for ContentType {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
ContentType::Text(text) => write!(f, "{}", text),
ContentType::MultiPart(multi_part) => {
let text_parts: Vec<String> = multi_part
.iter()
.filter_map(|part| {
if part.content_type == MultiPartContentType::Text {
part.text.clone()
} else {
panic!("Unsupported content type: {:?}", part.content_type);
}
})
.collect();
let combined_text = text_parts.join("\n");
write!(f, "{}", combined_text)
}
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)] #[derive(Debug, Clone, Serialize, Deserialize)]
pub struct Message { pub struct Message {
pub role: String, pub role: String,
#[serde(skip_serializing_if = "Option::is_none")] #[serde(skip_serializing_if = "Option::is_none")]
pub content: Option<String>, pub content: Option<ContentType>,
#[serde(skip_serializing_if = "Option::is_none")] #[serde(skip_serializing_if = "Option::is_none")]
pub model: Option<String>, pub model: Option<String>,
@ -237,7 +280,7 @@ impl ChatCompletionsResponse {
choices: vec![Choice { choices: vec![Choice {
message: Message { message: Message {
role: ASSISTANT_ROLE.to_string(), role: ASSISTANT_ROLE.to_string(),
content: Some(message), content: Some(ContentType::Text(message)),
model: Some(ARCH_FC_MODEL_NAME.to_string()), model: Some(ARCH_FC_MODEL_NAME.to_string()),
tool_calls: None, tool_calls: None,
tool_call_id: None, tool_call_id: None,
@ -378,6 +421,8 @@ pub fn to_server_events(chunks: Vec<ChatCompletionStreamResponse>) -> String {
#[cfg(test)] #[cfg(test)]
mod test { mod test {
use crate::api::open_ai::{ChatCompletionsRequest, ContentType, MultiPartContentType};
use super::{ChatCompletionStreamResponseServerEvents, Message}; use super::{ChatCompletionStreamResponseServerEvents, Message};
use pretty_assertions::assert_eq; use pretty_assertions::assert_eq;
use std::collections::HashMap; use std::collections::HashMap;
@ -447,7 +492,9 @@ mod test {
model: "gpt-3.5-turbo".to_string(), model: "gpt-3.5-turbo".to_string(),
messages: vec![Message { messages: vec![Message {
role: "user".to_string(), role: "user".to_string(),
content: Some("What city do you want to know the weather for?".to_string()), content: Some(ContentType::Text(
"What city do you want to know the weather for?".to_string(),
)),
model: None, model: None,
tool_calls: None, tool_calls: None,
tool_call_id: None, tool_call_id: None,
@ -677,4 +724,109 @@ data: [DONE]
"Hello! How can I assist you today?" "Hello! How can I assist you today?"
); );
} }
#[test]
fn test_chat_completions_request() {
const CHAT_COMPLETIONS_REQUEST: &str = r#"
{
"model": "gpt-3.5-turbo",
"messages": [
{
"role": "user",
"content": "What city do you want to know the weather for?"
}
]
}"#;
let chat_completions_request: ChatCompletionsRequest =
serde_json::from_str(CHAT_COMPLETIONS_REQUEST).unwrap();
assert_eq!(chat_completions_request.model, "gpt-3.5-turbo");
assert_eq!(
chat_completions_request.messages[0].content,
Some(ContentType::Text(
"What city do you want to know the weather for?".to_string()
))
);
}
#[test]
fn test_chat_completions_request_text_type() {
const CHAT_COMPLETIONS_REQUEST: &str = r#"
{
"model": "gpt-3.5-turbo",
"messages": [
{
"role": "user",
"content": [
{
"type": "text",
"text": "What city do you want to know the weather for?"
}
]
}
]
}
"#;
let chat_completions_request: ChatCompletionsRequest =
serde_json::from_str(CHAT_COMPLETIONS_REQUEST).unwrap();
assert_eq!(chat_completions_request.model, "gpt-3.5-turbo");
if let Some(ContentType::MultiPart(multi_part_content)) =
chat_completions_request.messages[0].content.as_ref()
{
assert_eq!(multi_part_content[0].content_type, MultiPartContentType::Text);
assert_eq!(
multi_part_content[0].text,
Some("What city do you want to know the weather for?".to_string())
);
} else {
panic!("Expected MultiPartContent");
}
}
#[test]
fn test_chat_completions_request_text_type_array() {
const CHAT_COMPLETIONS_REQUEST: &str = r#"
{
"model": "gpt-3.5-turbo",
"messages": [
{
"role": "user",
"content": [
{
"type": "text",
"text": "What city do you want to know the weather for?"
},
{
"type": "text",
"text": "hello world"
}
]
}
]
}
"#;
let chat_completions_request: ChatCompletionsRequest =
serde_json::from_str(CHAT_COMPLETIONS_REQUEST).unwrap();
assert_eq!(chat_completions_request.model, "gpt-3.5-turbo");
if let Some(ContentType::MultiPart(multi_part_content)) =
chat_completions_request.messages[0].content.as_ref()
{
assert_eq!(multi_part_content.len(), 2);
assert_eq!(multi_part_content[0].content_type, MultiPartContentType::Text);
assert_eq!(
multi_part_content[0].text,
Some("What city do you want to know the weather for?".to_string())
);
assert_eq!(multi_part_content[1].content_type, MultiPartContentType::Text);
assert_eq!(
multi_part_content[1].text,
Some("hello world".to_string())
);
} else {
panic!("Expected MultiPartContent");
}
}
} }

View file

@ -1,7 +1,7 @@
use crate::metrics::Metrics; use crate::metrics::Metrics;
use common::api::open_ai::{ use common::api::open_ai::{
ChatCompletionStreamResponseServerEvents, ChatCompletionsRequest, ChatCompletionsResponse, ChatCompletionStreamResponseServerEvents, ChatCompletionsRequest, ChatCompletionsResponse,
Message, StreamOptions, ContentType, Message, StreamOptions,
}; };
use common::configuration::{LlmProvider, LlmProviderType, Overrides}; use common::configuration::{LlmProvider, LlmProviderType, Overrides};
use common::consts::{ use common::consts::{
@ -384,7 +384,12 @@ impl HttpContext for StreamContext {
.messages .messages
.iter() .iter()
.fold(String::new(), |acc, m| { .fold(String::new(), |acc, m| {
acc + " " + m.content.as_ref().unwrap_or(&String::new()) acc + " "
+ m.content
.as_ref()
.unwrap_or(&ContentType::Text(String::new()))
.to_string()
.as_str()
}); });
// enforce ratelimits on ingress // enforce ratelimits on ingress
if let Err(e) = self.enforce_ratelimits(&deserialized_body.model, input_tokens_str.as_str()) if let Err(e) = self.enforce_ratelimits(&deserialized_body.model, input_tokens_str.as_str())

View file

@ -2,6 +2,7 @@ use crate::stream_context::{ResponseHandlerType, StreamCallContext, StreamContex
use common::{ use common::{
api::open_ai::{ api::open_ai::{
self, ArchState, ChatCompletionStreamResponse, ChatCompletionTool, ChatCompletionsRequest, self, ArchState, ChatCompletionStreamResponse, ChatCompletionTool, ChatCompletionsRequest,
ContentType,
}, },
consts::{ consts::{
ARCH_FC_MODEL_NAME, ARCH_INTERNAL_CLUSTER_NAME, ARCH_ROUTING_HEADER, ARCH_FC_MODEL_NAME, ARCH_INTERNAL_CLUSTER_NAME, ARCH_ROUTING_HEADER,
@ -237,22 +238,32 @@ impl HttpContext for StreamContext {
Duration::from_secs(5), Duration::from_secs(5),
); );
let call_context = StreamCallContext { if let Some(ContentType::Text(content)) =
response_handler_type: ResponseHandlerType::ArchFC, self.user_prompt.as_ref().unwrap().content.as_ref()
user_message: self.user_prompt.as_ref().unwrap().content.clone(), {
prompt_target_name: None, let call_context = StreamCallContext {
request_body: self.chat_completions_request.as_ref().unwrap().clone(), response_handler_type: ResponseHandlerType::ArchFC,
similarity_scores: None, user_message: Some(content.clone()),
upstream_cluster: Some(ARCH_INTERNAL_CLUSTER_NAME.to_string()), prompt_target_name: None,
upstream_cluster_path: Some("/function_calling".to_string()), request_body: self.chat_completions_request.as_ref().unwrap().clone(),
}; similarity_scores: None,
upstream_cluster: Some(ARCH_INTERNAL_CLUSTER_NAME.to_string()),
upstream_cluster_path: Some("/function_calling".to_string()),
};
if let Err(e) = self.http_call(call_args, call_context) { if let Err(e) = self.http_call(call_args, call_context) {
warn!("http_call failed: {:?}", e); warn!("http_call failed: {:?}", e);
self.send_server_error(ServerError::HttpDispatch(e), None); self.send_server_error(ServerError::HttpDispatch(e), None);
}
} else {
warn!("No content in the last user prompt");
self.send_server_error(
ServerError::LogicError("No content in the last user prompt".to_string()),
None,
);
} }
Action::Pause Action::Pause
} }
fn on_http_response_headers(&mut self, _num_headers: usize, _end_of_stream: bool) -> Action { fn on_http_response_headers(&mut self, _num_headers: usize, _end_of_stream: bool) -> Action {

View file

@ -2,7 +2,7 @@ use crate::metrics::Metrics;
use crate::tools::compute_request_path_body; use crate::tools::compute_request_path_body;
use common::api::open_ai::{ use common::api::open_ai::{
to_server_events, ArchState, ChatCompletionStreamResponse, ChatCompletionsRequest, to_server_events, ArchState, ChatCompletionStreamResponse, ChatCompletionsRequest,
ChatCompletionsResponse, Message, ToolCall, ChatCompletionsResponse, ContentType, Message, ToolCall,
}; };
use common::configuration::{Endpoint, Overrides, PromptTarget, Tracing}; use common::configuration::{Endpoint, Overrides, PromptTarget, Tracing};
use common::consts::{ use common::consts::{
@ -215,7 +215,7 @@ impl StreamContext {
Some(system_prompt) => { Some(system_prompt) => {
let system_prompt_message = Message { let system_prompt_message = Message {
role: SYSTEM_ROLE.to_string(), role: SYSTEM_ROLE.to_string(),
content: Some(system_prompt.clone()), content: Some(ContentType::Text(system_prompt.clone())),
model: None, model: None,
tool_calls: None, tool_calls: None,
tool_call_id: None, tool_call_id: None,
@ -279,6 +279,13 @@ impl StreamContext {
//TODO: add resolver name to the response so the client can send the response back to the correct resolver //TODO: add resolver name to the response so the client can send the response back to the correct resolver
let direct_response_str = if self.streaming_response { let direct_response_str = if self.streaming_response {
let content = model_server_response.choices[0]
.message
.content
.as_ref()
.unwrap()
.clone();
let chunks = vec![ let chunks = vec![
ChatCompletionStreamResponse::new( ChatCompletionStreamResponse::new(
self.arch_fc_response.clone(), self.arch_fc_response.clone(),
@ -287,14 +294,7 @@ impl StreamContext {
None, None,
), ),
ChatCompletionStreamResponse::new( ChatCompletionStreamResponse::new(
Some( Some(content.to_string()),
model_server_response.choices[0]
.message
.content
.as_ref()
.unwrap()
.clone(),
),
None, None,
Some(format!("{}-Chat", ARCH_FC_MODEL_NAME.to_owned())), Some(format!("{}-Chat", ARCH_FC_MODEL_NAME.to_owned())),
None, None,
@ -542,7 +542,7 @@ impl StreamContext {
messages.push({ messages.push({
Message { Message {
role: USER_ROLE.to_string(), role: USER_ROLE.to_string(),
content: Some(final_prompt), content: Some(ContentType::Text(final_prompt)),
model: None, model: None,
tool_calls: None, tool_calls: None,
tool_call_id: None, tool_call_id: None,
@ -612,7 +612,7 @@ impl StreamContext {
if system_prompt.is_some() { if system_prompt.is_some() {
let system_prompt_message = Message { let system_prompt_message = Message {
role: SYSTEM_ROLE.to_string(), role: SYSTEM_ROLE.to_string(),
content: system_prompt, content: Some(ContentType::Text(system_prompt.unwrap())),
model: None, model: None,
tool_calls: None, tool_calls: None,
tool_call_id: None, tool_call_id: None,
@ -639,7 +639,9 @@ impl StreamContext {
} else { } else {
Message { Message {
role: ASSISTANT_ROLE.to_string(), role: ASSISTANT_ROLE.to_string(),
content: self.arch_fc_response.as_ref().cloned(), content: Some(ContentType::Text(
self.arch_fc_response.as_ref().unwrap().clone(),
)),
model: Some(ARCH_FC_MODEL_NAME.to_string()), model: Some(ARCH_FC_MODEL_NAME.to_string()),
tool_calls: None, tool_calls: None,
tool_call_id: None, tool_call_id: None,
@ -650,7 +652,9 @@ impl StreamContext {
pub fn generate_api_response_message(&mut self) -> Message { pub fn generate_api_response_message(&mut self) -> Message {
Message { Message {
role: TOOL_ROLE.to_string(), role: TOOL_ROLE.to_string(),
content: self.tool_call_response.clone(), content: Some(ContentType::Text(
self.tool_call_response.as_ref().unwrap().clone(),
)),
model: None, model: None,
tool_calls: None, tool_calls: None,
tool_call_id: Some(self.tool_calls.as_ref().unwrap()[0].id.clone()), tool_call_id: Some(self.tool_calls.as_ref().unwrap()[0].id.clone()),
@ -688,7 +692,14 @@ impl StreamContext {
None, None,
), ),
ChatCompletionStreamResponse::new( ChatCompletionStreamResponse::new(
chat_completion_response.choices[0].message.content.clone(), Some(
chat_completion_response.choices[0]
.message
.content
.as_ref()
.unwrap()
.to_string(),
),
None, None,
Some(chat_completion_response.model.clone()), Some(chat_completion_response.model.clone()),
None, None,
@ -727,7 +738,7 @@ impl StreamContext {
Some(system_prompt) => { Some(system_prompt) => {
let system_prompt_message = Message { let system_prompt_message = Message {
role: SYSTEM_ROLE.to_string(), role: SYSTEM_ROLE.to_string(),
content: Some(system_prompt.clone()), content: Some(ContentType::Text(system_prompt.clone())),
model: None, model: None,
tool_calls: None, tool_calls: None,
tool_call_id: None, tool_call_id: None,
@ -748,7 +759,7 @@ impl StreamContext {
let message = format!("{}\ncontext: {}", user_message.content.unwrap(), api_resp); let message = format!("{}\ncontext: {}", user_message.content.unwrap(), api_resp);
messages.push(Message { messages.push(Message {
role: USER_ROLE.to_string(), role: USER_ROLE.to_string(),
content: Some(message), content: Some(ContentType::Text(message)),
model: None, model: None,
tool_calls: None, tool_calls: None,
tool_call_id: None, tool_call_id: None,
@ -781,7 +792,7 @@ fn check_intent_matched(model_server_response: &ChatCompletionsResponse) -> bool
.first() .first()
.and_then(|choice| choice.message.content.as_ref()); .and_then(|choice| choice.message.content.as_ref());
let content_has_value = content.is_some() && !content.unwrap().is_empty(); let content_has_value = content.is_some() && !content.unwrap().to_string().is_empty();
let tool_calls = model_server_response let tool_calls = model_server_response
.choices .choices
@ -807,7 +818,7 @@ impl Client for StreamContext {
#[cfg(test)] #[cfg(test)]
mod test { mod test {
use common::api::open_ai::{ChatCompletionsResponse, Choice, Message, ToolCall}; use common::api::open_ai::{ChatCompletionsResponse, Choice, ContentType, Message, ToolCall};
use crate::stream_context::check_intent_matched; use crate::stream_context::check_intent_matched;
@ -816,7 +827,7 @@ mod test {
let model_server_response = ChatCompletionsResponse { let model_server_response = ChatCompletionsResponse {
choices: vec![Choice { choices: vec![Choice {
message: Message { message: Message {
content: Some("".to_string()), content: Some(ContentType::Text("".to_string())),
tool_calls: Some(vec![]), tool_calls: Some(vec![]),
role: "assistant".to_string(), role: "assistant".to_string(),
model: None, model: None,
@ -835,7 +846,7 @@ mod test {
let model_server_response = ChatCompletionsResponse { let model_server_response = ChatCompletionsResponse {
choices: vec![Choice { choices: vec![Choice {
message: Message { message: Message {
content: Some("hello".to_string()), content: Some(ContentType::Text("hello".to_string())),
tool_calls: Some(vec![]), tool_calls: Some(vec![]),
role: "assistant".to_string(), role: "assistant".to_string(),
model: None, model: None,
@ -854,7 +865,7 @@ mod test {
let model_server_response = ChatCompletionsResponse { let model_server_response = ChatCompletionsResponse {
choices: vec![Choice { choices: vec![Choice {
message: Message { message: Message {
content: Some("".to_string()), content: Some(ContentType::Text("".to_string())),
tool_calls: Some(vec![ToolCall { tool_calls: Some(vec![ToolCall {
id: "1".to_string(), id: "1".to_string(),
function: common::api::open_ai::FunctionCallDetail { function: common::api::open_ai::FunctionCallDetail {

View file

@ -1,5 +1,5 @@
use common::api::open_ai::{ use common::api::open_ai::{
ChatCompletionsResponse, Choice, FunctionCallDetail, Message, ToolCall, ToolType, Usage, ChatCompletionsResponse, Choice, ContentType, FunctionCallDetail, Message, ToolCall, ToolType, Usage
}; };
use common::configuration::Configuration; use common::configuration::Configuration;
use http::StatusCode; use http::StatusCode;
@ -431,7 +431,7 @@ fn prompt_gateway_request_to_llm_gateway() {
index: Some(0), index: Some(0),
message: Message { message: Message {
role: "assistant".to_string(), role: "assistant".to_string(),
content: Some("hello from fake llm gateway".to_string()), content: Some(ContentType::Text("hello from fake llm gateway".to_string())),
model: None, model: None,
tool_calls: None, tool_calls: None,
tool_call_id: None, tool_call_id: None,

View file

@ -14,13 +14,13 @@ llm_providers:
- name: archgw-v1-router-model - name: archgw-v1-router-model
provider_interface: openai provider_interface: openai
model: cotran2/llama-1b-4-26 model: cotran2/llama-4-epoch
base_url: http://35.192.87.187:8000/v1 base_url: http://34.46.85.85:8000/v1
- name: gpt-4o-mini - name: gpt-4o
provider_interface: openai provider_interface: openai
access_key: $OPENAI_API_KEY access_key: $OPENAI_API_KEY
model: gpt-4o-mini model: gpt-4o
default: true default: true
- name: gpt-4o - name: gpt-4o
@ -35,5 +35,17 @@ llm_providers:
model: o4-mini model: o4-mini
usage: Requesting topic ideas specifically related to personal finance and budgeting. usage: Requesting topic ideas specifically related to personal finance and budgeting.
- name: code_generation
provider_interface: openai
access_key: $OPENAI_API_KEY
model: gpt-4.1
usage: Generating new code snippets, functions, or boilerplate based on user prompts or requirements
- name: code_understanding
provider_interface: openai
access_key: $OPENAI_API_KEY
model: gpt-4.1
usage: understand and explain existing code snippets, functions, or libraries
tracing: tracing:
random_sampling: 100 random_sampling: 100