magi-code 0.77.1

Repository-aware CLI coding agent for terminal work
Documentation
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)})
}

// Parse every object through a visitor so duplicate payload keys cannot change mutation intent.
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)
}