use std::sync::Arc;
use axum::{
Json, Router,
body::Bytes,
extract::State,
http::{HeaderMap, StatusCode, header},
response::{
IntoResponse, Response,
sse::{Event, KeepAlive, Sse},
},
routing::{get, post},
};
use futures::StreamExt as _;
use polyc_rpc_client::Sensitive;
use polyc_runtime::admission::AdmissionGate;
use serde_json::Value;
use subtle::ConstantTimeEq;
use crate::{
rpc,
store::TaskStore,
task::{ApprovalResponder, TurnRunner},
};
#[derive(Clone)]
pub struct AppState {
pub card: Arc<Value>,
pub runner: Arc<dyn TurnRunner>,
pub approvals: Arc<dyn ApprovalResponder>,
pub store: Arc<dyn TaskStore>,
pub turn_limit: AdmissionGate,
pub peers: PeerAuthenticator,
}
#[derive(Clone, Default)]
pub struct PeerAuthenticator {
entries: Arc<Vec<PeerCredential>>,
}
struct PeerCredential {
peer_id: String,
token: Sensitive<String>,
}
#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
pub enum PeerCredentialError {
#[error("A2A peer credentials must use peer-id=token entries")]
Malformed,
#[error("A2A peer credential ids and tokens must not be empty")]
Empty,
#[error("each A2A peer id must appear once")]
DuplicatePeer,
#[error("each A2A bearer token must identify exactly one peer")]
DuplicateToken,
}
impl PeerAuthenticator {
pub fn parse(raw: &str) -> Result<Self, PeerCredentialError> {
let mut entries: Vec<PeerCredential> = Vec::new();
for item in raw.split_whitespace() {
let (peer_id, token) = item.split_once('=').ok_or(PeerCredentialError::Malformed)?;
if peer_id.is_empty() || token.is_empty() {
return Err(PeerCredentialError::Empty);
}
if entries.iter().any(|entry| entry.peer_id == peer_id) {
return Err(PeerCredentialError::DuplicatePeer);
}
if entries.iter().any(|entry| entry.token.expose() == token) {
return Err(PeerCredentialError::DuplicateToken);
}
entries.push(PeerCredential {
peer_id: peer_id.to_owned(),
token: Sensitive::new(token.to_owned()),
});
}
Ok(Self {
entries: Arc::new(entries),
})
}
pub fn single(peer_id: &str, token: &str) -> Result<Self, PeerCredentialError> {
Self::parse(&format!("{peer_id}={token}"))
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.entries.is_empty()
}
fn authenticate(&self, provided: &str) -> Option<String> {
let mut authenticated = None;
for entry in self.entries.iter() {
let matches: bool = provided
.as_bytes()
.ct_eq(entry.token.expose().as_bytes())
.into();
if matches {
authenticated = Some(entry.peer_id.clone());
}
}
authenticated
}
}
pub fn router(state: AppState) -> Router {
Router::new()
.route("/.well-known/agent-card.json", get(agent_card))
.route("/", post(json_rpc))
.with_state(state)
}
async fn agent_card(State(state): State<AppState>) -> impl IntoResponse {
(
[(header::CONTENT_TYPE, "application/json")],
state.card.to_string(),
)
}
async fn json_rpc(State(state): State<AppState>, headers: HeaderMap, body: Bytes) -> Response {
let peer_id = match authenticate(&state, &headers) {
Ok(peer_id) => peer_id,
Err(rejection) => return *rejection,
};
if rpc::is_streaming_method(&body) {
let mut stream = rpc::handle_streaming_for_peer(state, &body, peer_id);
if !rpc::requires_durable_marker(&body) {
let stream = stream.map(|value| {
Ok::<_, std::convert::Infallible>(Event::default().data(value.to_string()))
});
return Sse::new(stream)
.keep_alive(KeepAlive::default())
.into_response();
}
let mut before_receipt = Vec::new();
let mut received = false;
while let Some(value) = stream.next().await {
if rpc::is_durable_marker(&value) {
received = true;
break;
}
before_receipt.push(value);
}
if !received {
let refusal = before_receipt.pop().unwrap_or_else(|| {
serde_json::json!({
"jsonrpc": "2.0",
"id": null,
"error": { "code": -32603, "message": "ingress ended without a durable receipt" }
})
});
return Json(refusal).into_response();
}
let stream = futures::stream::iter(before_receipt)
.chain(stream)
.filter(|value| futures::future::ready(!rpc::is_durable_marker(value)))
.map(|value| {
Ok::<_, std::convert::Infallible>(Event::default().data(value.to_string()))
});
return Sse::new(stream)
.keep_alive(KeepAlive::default())
.into_response();
}
Json(rpc::handle_for_peer(&state, &body, &peer_id).await).into_response()
}
fn authenticate(state: &AppState, headers: &HeaderMap) -> Result<String, Box<Response>> {
if state.peers.is_empty() {
return Err(Box::new(
(
StatusCode::SERVICE_UNAVAILABLE,
"A2A peer credentials are not configured",
)
.into_response(),
));
}
let Some(provided) = headers
.get(header::AUTHORIZATION)
.and_then(|value| value.to_str().ok())
.and_then(|value| value.strip_prefix("Bearer "))
else {
return Err(Box::new(
(StatusCode::UNAUTHORIZED, "missing bearer token").into_response(),
));
};
state
.peers
.authenticate(provided)
.ok_or_else(|| Box::new((StatusCode::UNAUTHORIZED, "invalid bearer token").into_response()))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn peer_credentials_refuse_shared_tokens_and_duplicate_ids() {
assert_eq!(
PeerAuthenticator::parse("peer-a=one peer-b=one").err(),
Some(PeerCredentialError::DuplicateToken)
);
assert_eq!(
PeerAuthenticator::parse("peer-a=one peer-a=two").err(),
Some(PeerCredentialError::DuplicatePeer)
);
}
}