1use 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
15pub const DEFAULT_TIMEOUT: Duration = Duration::from_secs(5);
17pub const MAX_RESPONSE_BYTES: usize = 65_536;
19
20#[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#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
44#[non_exhaustive]
45pub enum HookConfigError {
46 #[error("a non-isolated hook channel requires a signer")]
49 SignerRequired,
50 #[error("an unsigned hook channel must not report an origin")]
52 UnsignedOrigin,
53 #[error("a signed hook channel must report its canonical origin")]
56 ChannelAudience,
57 #[error("{0} is not a canonical origin (plain HTTP is loopback-only)")]
60 Audience(&'static str),
61}
62
63fn 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
76struct 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
89pub(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#[derive(Debug, Clone, Copy, PartialEq, Eq)]
103pub(crate) struct CallFailure(pub(crate) &'static str);
104
105pub struct HookClient<C> {
108 channel: C,
109 signer: Option<HookSigner>,
110 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 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 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 #[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 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 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 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 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}