1use std::str::FromStr;
4
5use compact_str::CompactString;
6use serde::de::DeserializeOwned;
7use serde::{Deserialize, Serialize};
8
9use crate::network::ChainId;
10
11#[derive(Debug, Clone, Eq, PartialEq, Serialize, Deserialize)]
13#[serde(rename_all = "camelCase", deny_unknown_fields)]
14#[non_exhaustive]
15pub struct PaymentRequirements<
16 TScheme = CompactString,
17 TAmount = CompactString,
18 TAddress = CompactString,
19 TExtra = serde_json::Value,
20> {
21 pub scheme: TScheme,
23 pub network: ChainId,
25 pub amount: TAmount,
27 pub pay_to: TAddress,
29 pub max_timeout_seconds: u64,
31 pub asset: TAddress,
33 #[serde(default = "Option::default", skip_serializing_if = "Option::is_none")]
35 pub extra: Option<TExtra>,
36}
37
38#[must_use]
40pub fn find_matching_requirements<'a>(
41 available: &'a [PaymentRequirements],
42 accepted: &PaymentRequirements,
43) -> Option<&'a PaymentRequirements> {
44 available
45 .iter()
46 .find(|req| req.matches_payload_accepted(accepted))
47}
48
49impl<TScheme, TAmount, TAddress, TExtra> PaymentRequirements<TScheme, TAmount, TAddress, TExtra> {
50 #[must_use]
52 pub const fn new(
53 scheme: TScheme,
54 network: ChainId,
55 amount: TAmount,
56 pay_to: TAddress,
57 asset: TAddress,
58 max_timeout_seconds: u64,
59 ) -> Self {
60 Self {
61 scheme,
62 network,
63 amount,
64 pay_to,
65 asset,
66 max_timeout_seconds,
67 extra: None,
68 }
69 }
70
71 #[must_use]
73 pub fn with_extra(mut self, extra: TExtra) -> Self {
74 self.extra = Some(extra);
75 self
76 }
77
78 #[must_use]
80 pub fn with_optional_extra(mut self, extra: Option<TExtra>) -> Self {
81 self.extra = extra;
82 self
83 }
84}
85
86impl<TScheme, TAmount, TAddress, TExtra> PaymentRequirements<TScheme, TAmount, TAddress, TExtra>
87where
88 TScheme: PartialEq,
89 TAmount: PartialEq,
90 TAddress: PartialEq,
91 TExtra: Serialize,
92{
93 #[must_use]
96 pub fn matches_payload_accepted(&self, accepted: &Self) -> bool {
97 self.matches_payload_accepted_with_dynamic(accepted, &[])
98 }
99
100 #[must_use]
103 pub fn matches_payload_accepted_with_dynamic(
104 &self,
105 accepted: &Self,
106 dynamic_extra_fields: &[&str],
107 ) -> bool {
108 self.scheme == accepted.scheme
109 && self.network == accepted.network
110 && self.amount == accepted.amount
111 && self.asset == accepted.asset
112 && self.pay_to == accepted.pay_to
113 && self.max_timeout_seconds == accepted.max_timeout_seconds
114 && extra_contains_subset(
115 self.extra.as_ref(),
116 accepted.extra.as_ref(),
117 dynamic_extra_fields,
118 )
119 }
120}
121
122fn extra_contains_subset<T: Serialize>(
123 required: Option<&T>,
124 accepted: Option<&T>,
125 dynamic_extra_fields: &[&str],
126) -> bool {
127 let Some(required_extra) = required else {
128 return true;
129 };
130 let Ok(required_value) = serde_json::to_value(required_extra) else {
131 return false;
132 };
133 let accepted_value = accepted.and_then(|extra| serde_json::to_value(extra).ok());
134 if accepted.is_some() && accepted_value.is_none() {
135 return false;
136 }
137 let required_omitted = omit_fields(&required_value, dynamic_extra_fields);
138 let accepted_omitted = accepted_value
139 .as_ref()
140 .map(|value| omit_fields(value, dynamic_extra_fields));
141 object_contains_subset(&required_omitted, accepted_omitted.as_ref())
142}
143
144fn omit_fields(value: &serde_json::Value, fields: &[&str]) -> serde_json::Value {
145 if fields.is_empty() {
146 return value.clone();
147 }
148 let serde_json::Value::Object(map) = value else {
149 return value.clone();
150 };
151 let mut copied = map.clone();
152 for field in fields {
153 copied.remove(*field);
154 }
155 serde_json::Value::Object(copied)
156}
157
158fn object_contains_subset(
160 expected: &serde_json::Value,
161 actual: Option<&serde_json::Value>,
162) -> bool {
163 let serde_json::Value::Object(expected_map) = expected else {
164 return actual.is_some_and(|got| got == expected);
165 };
166 let Some(serde_json::Value::Object(actual_map)) = actual else {
167 return false;
168 };
169 expected_map.iter().all(|(key, value)| {
170 actual_map.get(key).map_or_else(
171 || value.is_null(),
172 |got| object_contains_subset(value, Some(got)),
173 )
174 })
175}
176
177impl PaymentRequirements {
178 #[must_use]
182 pub fn as_concrete<TScheme, TAmount, TAddress, TExtra>(
183 &self,
184 ) -> Option<PaymentRequirements<TScheme, TAmount, TAddress, TExtra>>
185 where
186 TScheme: FromStr,
187 TAmount: FromStr,
188 TAddress: FromStr,
189 TExtra: DeserializeOwned,
190 {
191 let scheme = self.scheme.parse::<TScheme>().ok()?;
192 let amount = self.amount.parse::<TAmount>().ok()?;
193 let pay_to = self.pay_to.parse::<TAddress>().ok()?;
194 let asset = self.asset.parse::<TAddress>().ok()?;
195 let extra = self
196 .extra
197 .as_ref()
198 .and_then(|v| serde_json::from_value(v.clone()).ok());
199 Some(PaymentRequirements {
200 scheme,
201 network: self.network.clone(),
202 amount,
203 pay_to,
204 max_timeout_seconds: self.max_timeout_seconds,
205 asset,
206 extra,
207 })
208 }
209}
210
211#[cfg(test)]
212#[allow(clippy::unwrap_used, reason = "unit tests panic on assertion failure")]
213mod tests {
214 use super::*;
215
216 #[test]
217 fn rejects_unknown_top_level_field() {
218 let json = serde_json::json!({
219 "scheme": "exact",
220 "network": "eip155:8453",
221 "amount": "1",
222 "payTo": "0x0",
223 "maxTimeoutSeconds": 60,
224 "asset": "0x0",
225 "unknownField": 1
226 });
227 assert!(serde_json::from_value::<PaymentRequirements>(json).is_err());
228 }
229
230 fn sample(network: &str, amount: &str, pay_to: &str, timeout: u64) -> PaymentRequirements {
231 PaymentRequirements::new(
232 "exact".into(),
233 network.parse().unwrap(),
234 amount.into(),
235 pay_to.into(),
236 "USDC".into(),
237 timeout,
238 )
239 }
240
241 #[test]
242 fn find_matching_requirements_go_semantics() {
243 let a = sample("eip155:1", "1000000", "0xrecipient1", 60);
244 let b = sample("eip155:8453", "2000000", "0xrecipient2", 30);
245 let available = [a.clone(), b.clone()];
246
247 let matched = find_matching_requirements(&available, &b).unwrap();
248 assert_eq!(matched.network.to_string(), "eip155:8453");
249 assert_eq!(matched.max_timeout_seconds, 30);
250
251 let mut timeout_miss = b;
252 timeout_miss.max_timeout_seconds = 999;
253 assert!(find_matching_requirements(&available, &timeout_miss).is_none());
254
255 let mut miss = a;
256 miss.scheme = "nonexistent".into();
257 assert!(find_matching_requirements(&available, &miss).is_none());
258 }
259
260 #[test]
261 fn matches_when_accepted_extra_has_additional_object_keys() {
262 let required =
263 sample("eip155:8453", "1000000", "0xabc", 300).with_extra(serde_json::json!({
264 "name": "USDC",
265 "version": "2",
266 "nested": { "required": true }
267 }));
268 let accepted =
269 sample("eip155:8453", "1000000", "0xabc", 300).with_extra(serde_json::json!({
270 "name": "USDC",
271 "version": "2",
272 "nested": { "required": true, "clientOnly": "ok" },
273 "channelState": { "chargedCumulativeAmount": "2000" }
274 }));
275 assert!(required.matches_payload_accepted(&accepted));
276 }
277
278 #[test]
279 fn matches_when_required_extra_is_absent() {
280 let required = sample("eip155:8453", "1000000", "0xabc", 300);
281 let accepted = sample("eip155:8453", "1000000", "0xabc", 300)
282 .with_extra(serde_json::json!({ "clientOnly": true }));
283 assert!(required.matches_payload_accepted(&accepted));
284 }
285
286 #[test]
287 fn matches_when_required_null_key_is_missing_on_accepted() {
288 let required = sample("eip155:8453", "1000000", "0xabc", 300)
289 .with_extra(serde_json::json!({ "k": null }));
290 let accepted =
291 sample("eip155:8453", "1000000", "0xabc", 300).with_extra(serde_json::json!({}));
292 assert!(required.matches_payload_accepted(&accepted));
293 }
294
295 #[test]
296 fn does_not_match_when_accepted_extra_overwrites_server_field() {
297 let required = sample("eip155:8453", "1000000", "0xabc", 300)
298 .with_extra(serde_json::json!({ "name": "USDC", "version": "2" }));
299 let accepted = sample("eip155:8453", "1000000", "0xabc", 300)
300 .with_extra(serde_json::json!({ "name": "USDC", "version": "3" }));
301 assert!(!required.matches_payload_accepted(&accepted));
302 }
303
304 #[test]
305 fn does_not_match_when_accepted_extra_array_is_a_superset() {
306 let required = sample("eip155:8453", "1000000", "0xabc", 300)
307 .with_extra(serde_json::json!({ "allowedSigners": ["0xalice"] }));
308 let accepted = sample("eip155:8453", "1000000", "0xabc", 300)
309 .with_extra(serde_json::json!({ "allowedSigners": ["0xmallory", "0xalice"] }));
310 assert!(!required.matches_payload_accepted(&accepted));
311 }
312
313 #[test]
314 fn does_not_match_when_accepted_extra_array_is_reordered() {
315 let required = sample("eip155:8453", "1000000", "0xabc", 300)
316 .with_extra(serde_json::json!({ "allowedSigners": ["0xalice", "0xbob"] }));
317 let accepted = sample("eip155:8453", "1000000", "0xabc", 300)
318 .with_extra(serde_json::json!({ "allowedSigners": ["0xbob", "0xalice"] }));
319 assert!(!required.matches_payload_accepted(&accepted));
320 }
321
322 #[test]
323 fn does_not_match_when_accepted_extra_omits_server_fields() {
324 let required = sample("eip155:8453", "1000000", "0xabc", 300)
325 .with_extra(serde_json::json!({ "name": "USDC", "version": "2" }));
326 let accepted = sample("eip155:8453", "1000000", "0xabc", 300)
327 .with_extra(serde_json::json!({ "name": "USDC" }));
328 assert!(!required.matches_payload_accepted(&accepted));
329 }
330
331 #[test]
332 fn matches_when_only_declared_dynamic_extra_fields_differ() {
333 let required =
334 sample("solana:mainnet", "1000000", "PayTo1", 60).with_extra(serde_json::json!({
335 "feePayer": "FeePayer111111111111111111111111111111111",
336 "recentBlockhash": "freshBlockhash",
337 "lastValidBlockHeight": "200"
338 }));
339 let accepted =
340 sample("solana:mainnet", "1000000", "PayTo1", 60).with_extra(serde_json::json!({
341 "feePayer": "FeePayer111111111111111111111111111111111",
342 "recentBlockhash": "staleBlockhash",
343 "lastValidBlockHeight": "100"
344 }));
345 assert!(required.matches_payload_accepted_with_dynamic(
346 &accepted,
347 &["recentBlockhash", "lastValidBlockHeight"],
348 ));
349 assert!(!required.matches_payload_accepted(&accepted));
350 }
351
352 #[test]
353 fn does_not_match_when_static_extra_differs_despite_dynamic_fields() {
354 let required =
355 sample("solana:mainnet", "1000000", "PayTo1", 60).with_extra(serde_json::json!({
356 "feePayer": "FeePayer111111111111111111111111111111111",
357 "recentBlockhash": "freshBlockhash"
358 }));
359 let accepted =
360 sample("solana:mainnet", "1000000", "PayTo1", 60).with_extra(serde_json::json!({
361 "feePayer": "OtherPayer1111111111111111111111111111111",
362 "recentBlockhash": "staleBlockhash"
363 }));
364 assert!(!required.matches_payload_accepted_with_dynamic(
365 &accepted,
366 &["recentBlockhash", "lastValidBlockHeight"],
367 ));
368 }
369}