use crate::service::protocol::{self, ServiceErrorCode};
use serde::{Deserialize, Serialize};
use serde_json::{Value, json};
pub(super) const MAX_COUNTER: u64 = 9_007_199_254_740_991;
pub(super) const OPERATIONS: [&str; 22] = [
"initialize",
"status",
"capabilities",
"session.list",
"session.create",
"session.active",
"session.claim",
"session.detach",
"session.replay",
"turn.start",
"turn.cancel",
"operation.lookup",
"auth.status",
"auth.login.start",
"auth.login.callback",
"auth.login.cancel",
"auth.logout",
"catalog.providers",
"catalog.models",
"catalog.refresh",
"config.get",
"config.set",
];
#[derive(Clone, Debug, Deserialize, Serialize, PartialEq, Eq)]
#[serde(deny_unknown_fields)]
pub(super) struct Control {
pub(super) grant_id: String,
pub(super) generation: u64,
}
#[derive(Clone, Debug, Deserialize, Serialize)]
#[serde(deny_unknown_fields)]
pub(super) struct Request {
pub(super) protocol_version: u16,
pub(super) kind: String,
pub(super) request_id: String,
pub(super) instance_id: Option<String>,
pub(super) connection_id: Option<String>,
pub(super) session_id: Option<String>,
pub(super) operation_id: Option<String>,
pub(super) control: Option<Control>,
pub(super) method: String,
pub(super) payload: Value,
}
impl Request {
pub(super) fn mutation(&self) -> bool {
matches!(
self.method.as_str(),
"session.create"
| "session.claim"
| "session.detach"
| "turn.start"
| "turn.cancel"
| "auth.login.start"
| "auth.login.callback"
| "auth.login.cancel"
| "auth.logout"
| "catalog.refresh"
| "config.set"
)
}
pub(super) fn needs_control(&self) -> bool {
matches!(
self.method.as_str(),
"session.detach" | "turn.start" | "turn.cancel"
)
}
pub(super) fn validate(&self) -> Result<(), ServiceErrorCode> {
if self.protocol_version != 2 {
return Err(ServiceErrorCode::UnsupportedVersion);
}
if self.kind != "request"
|| !valid_id(&self.request_id)
|| self.method.len() > protocol::MAX_NAME_BYTES
|| [
&self.instance_id,
&self.connection_id,
&self.session_id,
&self.operation_id,
]
.iter()
.any(|id| id.as_ref().is_some_and(|id| !valid_id(id)))
{
return Err(ServiceErrorCode::InvalidRequest);
}
protocol::payload_is_bounded(&self.payload)?;
if !OPERATIONS.contains(&self.method.as_str()) {
return Err(ServiceErrorCode::UnsupportedOperation);
}
let needs_session = self.needs_control()
|| matches!(self.method.as_str(), "session.claim" | "session.replay");
if needs_session != self.session_id.is_some()
|| self.needs_control() != self.control.is_some()
|| self.mutation() != self.operation_id.is_some()
{
return Err(ServiceErrorCode::InvalidPayload);
}
if self.control.as_ref().is_some_and(|control| {
!valid_id(&control.grant_id)
|| control.generation == 0
|| control.generation > MAX_COUNTER
}) {
return Err(ServiceErrorCode::InvalidPayload);
}
if self.method == "turn.start" {
#[derive(Deserialize)]
#[serde(deny_unknown_fields)]
struct TurnPrompt {
prompt: String,
}
let params: TurnPrompt = serde_json::from_value(self.payload.clone())
.map_err(|_| ServiceErrorCode::InvalidPayload)?;
if params.prompt.trim().is_empty() {
return Err(ServiceErrorCode::InvalidPayload);
}
}
Ok(())
}
pub(super) fn legacy(&self, method: &str) -> protocol::ServiceRequest {
protocol::ServiceRequest {
protocol_version: 1,
kind: protocol::MessageKind::Request,
request_id: self.request_id.clone(),
session_id: self.session_id.clone(),
method: method.to_owned(),
payload: self.payload.clone(),
}
}
}
pub(super) fn valid_id(value: &str) -> bool {
!value.is_empty() && value.len() <= 128
}
pub(super) fn error(code: ServiceErrorCode) -> Value {
json!({"code":code,"message":code.message()})
}
pub(super) fn response(
instance: &str,
connection: &str,
request: &Request,
result: Result<Value, ServiceErrorCode>,
) -> Value {
let (payload, error) = match result {
Ok(payload) => (payload, Value::Null),
Err(code) => (Value::Null, error(code)),
};
json!({"protocol_version":2,"kind":"response","instance_id":instance,"connection_id":connection,
"request_id":safe_text(&request.request_id, 128),
"session_id":request.session_id.as_deref().and_then(|id| safe_text(id, 128)),
"operation_id":request.operation_id.as_deref().and_then(|id| safe_text(id, 128)),
"method":safe_text(&request.method, 64),"payload":payload,"error":error})
}
fn safe_text(value: &str, maximum: usize) -> Option<&str> {
(!value.is_empty() && value.len() <= maximum).then_some(value)
}
pub(super) fn decoding_error(
instance: &str,
connection: &str,
bytes: &[u8],
code: ServiceErrorCode,
) -> Value {
let value = serde_json::from_slice::<UniqueJson>(bytes)
.ok()
.map(|value| value.0)
.unwrap_or(Value::Null);
let id = |field: &str, maximum| {
value[field]
.as_str()
.and_then(|value| safe_text(value, maximum))
};
json!({"protocol_version":2,"kind":"response","instance_id":instance,"connection_id":connection,
"request_id":id("request_id",128),"session_id":id("session_id",128),"operation_id":id("operation_id",128),
"method":id("method",64),"payload":null,"error":error(code)})
}
struct UniqueJson(Value);
impl<'de> Deserialize<'de> for UniqueJson {
fn deserialize<D: serde::Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
struct Visitor;
impl<'de> serde::de::Visitor<'de> for Visitor {
type Value = UniqueJson;
fn expecting(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str("JSON without duplicate keys")
}
fn visit_bool<E: serde::de::Error>(self, value: bool) -> Result<UniqueJson, E> {
Ok(UniqueJson(json!(value)))
}
fn visit_i64<E: serde::de::Error>(self, value: i64) -> Result<UniqueJson, E> {
Ok(UniqueJson(json!(value)))
}
fn visit_u64<E: serde::de::Error>(self, value: u64) -> Result<UniqueJson, E> {
Ok(UniqueJson(json!(value)))
}
fn visit_f64<E: serde::de::Error>(self, value: f64) -> Result<UniqueJson, E> {
Ok(UniqueJson(json!(value)))
}
fn visit_str<E: serde::de::Error>(self, value: &str) -> Result<UniqueJson, E> {
Ok(UniqueJson(json!(value)))
}
fn visit_unit<E: serde::de::Error>(self) -> Result<UniqueJson, E> {
Ok(UniqueJson(Value::Null))
}
fn visit_seq<A: serde::de::SeqAccess<'de>>(
self,
mut seq: A,
) -> Result<UniqueJson, A::Error> {
let mut values = Vec::new();
while let Some(UniqueJson(value)) = seq.next_element()? {
values.push(value);
}
Ok(UniqueJson(Value::Array(values)))
}
fn visit_map<A: serde::de::MapAccess<'de>>(
self,
mut map: A,
) -> Result<UniqueJson, A::Error> {
let mut values = serde_json::Map::new();
while let Some((key, UniqueJson(value))) = map.next_entry::<String, UniqueJson>()? {
if values.insert(key, value).is_some() {
return Err(serde::de::Error::custom("duplicate key"));
}
}
Ok(UniqueJson(Value::Object(values)))
}
}
deserializer.deserialize_any(Visitor)
}
}
pub(super) fn decode(bytes: &[u8]) -> Result<Request, ServiceErrorCode> {
if bytes.len() > protocol::MAX_RECORD_BYTES {
return Err(ServiceErrorCode::RecordTooLarge);
}
let UniqueJson(value) =
serde_json::from_slice(bytes).map_err(|_| ServiceErrorCode::InvalidJson)?;
for key in [
"instance_id",
"connection_id",
"session_id",
"operation_id",
"control",
] {
if value.get(key).is_none() {
return Err(ServiceErrorCode::InvalidRequest);
}
}
serde_json::from_value(value).map_err(|_| ServiceErrorCode::InvalidRequest)
}