Skip to main content

mkit_server/hooks/
client.rs

1//! One signed, bounded, size-checked hook call.
2
3use core::time::Duration;
4use std::sync::Arc;
5
6use mkit_core::write_auth::validate_audience;
7use serde::Serialize;
8use serde::de::DeserializeOwned;
9use zeroize::Zeroizing;
10
11use super::channel::{ChannelError, HookChannel, HookRequest, HookResponse};
12use super::{HookSigner, NonceSource, OsNonces};
13use crate::rt::{Clock, Sleep, with_timeout};
14
15/// The default per-call timeout (SPEC-SERVER §8, informative).
16pub const DEFAULT_TIMEOUT: Duration = Duration::from_secs(5);
17/// The largest response body core accepts (SPEC-SERVER §6.6).
18pub const MAX_RESPONSE_BYTES: usize = 65_536;
19
20/// A hook RPC this adapter calls.
21#[derive(Debug, Clone, Copy, PartialEq, Eq)]
22pub(crate) enum Rpc {
23    Authorize,
24    Admit,
25    Outcome,
26    CachePurge,
27    Inspect,
28}
29
30impl Rpc {
31    pub(crate) const fn path(self) -> &'static str {
32        match self {
33            Self::Authorize => "/mkit.server.hooks.v1.HooksService/Authorize",
34            Self::Admit => "/mkit.server.hooks.v1.HooksService/Admit",
35            Self::Outcome => "/mkit.server.hooks.v1.HooksService/Outcome",
36            Self::CachePurge => "/mkit.server.hooks.v1.HooksService/CachePurge",
37            Self::Inspect => "/mkit.server.hooks.v1.HooksService/Inspect",
38        }
39    }
40}
41
42/// A refused adapter configuration.
43#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
44#[non_exhaustive]
45pub enum HookConfigError {
46    /// A channel that is not an isolated service binding must sign
47    /// (SPEC-SERVER §7.1 versus §7.3).
48    #[error("a non-isolated hook channel requires a signer")]
49    SignerRequired,
50    /// An unsigned channel is a service binding, which has no origin.
51    #[error("an unsigned hook channel must not report an origin")]
52    UnsignedOrigin,
53    /// Signing binds the hook endpoint's canonical origin, so the channel
54    /// must name one.
55    #[error("a signed hook channel must report its canonical origin")]
56    ChannelAudience,
57    /// The named origin is not a canonical HTTP(S) origin, or is plain HTTP
58    /// to a host that is not loopback (SPEC-SERVER §6.1).
59    #[error("{0} is not a canonical origin (plain HTTP is loopback-only)")]
60    Audience(&'static str),
61}
62
63/// Whether the authority of a canonical `http://` origin is a loopback host.
64fn loopback(authority: &str) -> bool {
65    let host = match authority.strip_prefix('[') {
66        Some(v6) => v6.split(']').next().unwrap_or_default(),
67        None => authority.split(':').next().unwrap_or_default(),
68    };
69    host == "localhost"
70        || host == "::1"
71        || host
72            .parse::<std::net::Ipv4Addr>()
73            .is_ok_and(|ip| ip.is_loopback())
74}
75
76/// Counts the bytes a value serialises to, without keeping any.
77struct Counter(usize);
78
79impl std::io::Write for Counter {
80    fn write(&mut self, bytes: &[u8]) -> std::io::Result<usize> {
81        self.0 += bytes.len();
82        Ok(bytes.len())
83    }
84    fn flush(&mut self) -> std::io::Result<()> {
85        Ok(())
86    }
87}
88
89/// Serialise into one exactly sized `Zeroizing` buffer. Growing a `Vec` while
90/// serialising would free unwiped partial copies of any credential inside.
91pub(super) fn encode<Req: Serialize>(request: &Req) -> Result<Zeroizing<Vec<u8>>, CallFailure> {
92    let fail = |_| CallFailure("request not encodable");
93    let mut size = Counter(0);
94    serde_json::to_writer(&mut size, request).map_err(fail)?;
95    let mut body = Zeroizing::new(Vec::with_capacity(size.0));
96    serde_json::to_writer(&mut *body, request).map_err(fail)?;
97    Ok(body)
98}
99
100/// Why a call produced no usable answer. The reason is a fixed string, never
101/// hook-controlled text, so it is safe to log.
102#[derive(Debug, Clone, Copy, PartialEq, Eq)]
103pub(crate) struct CallFailure(pub(crate) &'static str);
104
105/// The shared client the per-role types (`RemoteAuthorizer`,
106/// `RemoteAdmission`, `RemoteOutcomes`) hold behind an `Arc`.
107pub struct HookClient<C> {
108    channel: C,
109    signer: Option<HookSigner>,
110    /// The channel's origin, validated once at construction; `Some` exactly
111    /// when the client signs.
112    hook_audience: Option<String>,
113    server_audience: String,
114    clock: Arc<dyn Clock>,
115    sleep: Arc<dyn Sleep>,
116    nonces: Arc<dyn NonceSource>,
117}
118
119impl<C> core::fmt::Debug for HookClient<C> {
120    fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
121        f.debug_struct("HookClient")
122            .field("server_audience", &self.server_audience)
123            .field("signer", &self.signer)
124            .finish_non_exhaustive()
125    }
126}
127
128impl<C: HookChannel> HookClient<C> {
129    /// A client over `channel`. `server_audience` is this server's own
130    /// canonical origin (the `audience` hook bodies carry, not the hook
131    /// endpoint's). `clock` stamps signatures and `sleep` enforces call
132    /// timeouts.
133    ///
134    /// # Errors
135    /// [`HookConfigError`] when the channel is neither signed nor isolated,
136    /// or an origin is malformed.
137    pub fn new(
138        channel: C,
139        server_audience: impl Into<String>,
140        signer: Option<HookSigner>,
141        clock: Arc<dyn Clock>,
142        sleep: Arc<dyn Sleep>,
143    ) -> Result<Self, HookConfigError> {
144        let server_audience = server_audience.into();
145        validate_audience(&server_audience).map_err(|_| HookConfigError::Audience("server"))?;
146        if signer.is_none() && !channel.isolated() {
147            return Err(HookConfigError::SignerRequired);
148        }
149        // A signed channel must report its origin, since the signature binds
150        // it, and is held to §6.1; an unsigned one is a service binding and
151        // has none (§7.3).
152        match (channel.audience(), signer.is_some()) {
153            (Some(_), false) => return Err(HookConfigError::UnsignedOrigin),
154            (Some(origin), true) => {
155                let plain_remote = origin
156                    .strip_prefix("http://")
157                    .is_some_and(|rest| !loopback(rest));
158                if validate_audience(origin).is_err() || plain_remote {
159                    return Err(HookConfigError::Audience("hook"));
160                }
161            }
162            (None, true) => return Err(HookConfigError::ChannelAudience),
163            (None, false) => {}
164        }
165        let hook_audience = channel.audience().map(str::to_owned);
166        Ok(Self {
167            channel,
168            signer,
169            hook_audience,
170            server_audience,
171            clock,
172            sleep,
173            nonces: Arc::new(OsNonces),
174        })
175    }
176
177    /// Replace the nonce source. A constant source makes the hook reject every
178    /// request after the first as a replay, so this fails closed; it exists
179    /// for tests that pin a signature.
180    #[must_use]
181    pub fn with_nonce_source(mut self, nonces: Arc<dyn NonceSource>) -> Self {
182        self.nonces = nonces;
183        self
184    }
185
186    #[cfg(test)]
187    pub(crate) fn channel(&self) -> &C {
188        &self.channel
189    }
190
191    /// This server's canonical origin, as bodies carry it.
192    pub(crate) fn server_audience(&self) -> &str {
193        &self.server_audience
194    }
195
196    pub(crate) fn is_signed(&self) -> bool {
197        self.signer.is_some()
198    }
199
200    /// Sign and send `request` to `rpc` and return the hook's answer whatever
201    /// its status. Every signature carries a fresh nonce and validity window.
202    async fn send<Req: Serialize>(
203        &self,
204        rpc: Rpc,
205        request: &Req,
206        timeout: Duration,
207    ) -> Result<HookResponse, CallFailure> {
208        let body = encode(request)?;
209        let mut headers = vec![
210            ("Content-Type", "application/json".to_owned()),
211            ("Connect-Protocol-Version", "1".to_owned()),
212        ];
213        if let Some(signer) = &self.signer {
214            let audience = self
215                .hook_audience
216                .as_deref()
217                .ok_or(CallFailure("no hook origin"))?;
218            let mut nonce = [0u8; 32];
219            if !self.nonces.fill(&mut nonce) {
220                return Err(CallFailure("no nonce"));
221            }
222            headers.extend(
223                signer
224                    .headers(audience, rpc.path(), &body, self.clock.now_ms(), &nonce)
225                    .map_err(|_| CallFailure("cannot sign"))?,
226            );
227        }
228        let call = self.channel.call(HookRequest {
229            procedure: rpc.path(),
230            headers,
231            body,
232            timeout,
233            max_response_bytes: MAX_RESPONSE_BYTES,
234        });
235        with_timeout(&*self.sleep, timeout, call)
236            .await
237            .map_err(|_| CallFailure("timeout"))?
238            .map_err(|err| match err {
239                ChannelError::Timeout => CallFailure("timeout"),
240                ChannelError::TooLarge => CallFailure("response too large"),
241                _ => CallFailure("transport error"),
242            })
243    }
244
245    /// Call a decision RPC: a 2xx JSON answer of at most 64 KiB that decodes
246    /// as `Res`, or a failure (unknown JSON fields are ignored).
247    pub(crate) async fn decide<Req, Res>(
248        &self,
249        rpc: Rpc,
250        request: &Req,
251        timeout: Duration,
252    ) -> Result<Res, CallFailure>
253    where
254        Req: Serialize,
255        Res: DeserializeOwned,
256    {
257        let response = self.send(rpc, request, timeout).await?;
258        if !(200..300).contains(&response.status) {
259            return Err(CallFailure("non-2xx status"));
260        }
261        if response.body.len() > MAX_RESPONSE_BYTES {
262            return Err(CallFailure("response too large"));
263        }
264        let json = response.content_type.as_deref().is_some_and(|value| {
265            value
266                .split(';')
267                .next()
268                .is_some_and(|media| media.trim().eq_ignore_ascii_case("application/json"))
269        });
270        if !json {
271            return Err(CallFailure("non-JSON content type"));
272        }
273        serde_json::from_slice(&response.body).map_err(|_| CallFailure("malformed response"))
274    }
275
276    /// Call a delivery RPC: any 2xx acknowledges, whatever the body.
277    pub(crate) async fn deliver<Req: Serialize>(
278        &self,
279        rpc: Rpc,
280        request: &Req,
281        timeout: Duration,
282    ) -> Result<(), CallFailure> {
283        let response = self.send(rpc, request, timeout).await?;
284        if (200..300).contains(&response.status) {
285            Ok(())
286        } else {
287            Err(CallFailure("non-2xx status"))
288        }
289    }
290}