1use axum::extract::{FromRequestParts, Query};
2use axum::http::request::Parts;
3use serde::Deserialize;
4use toolkit_canonical_errors::CanonicalError;
5use toolkit_odata::errors::OdataError;
6use toolkit_odata::{CursorV1, Error as ODataError, ODataOrderBy, OrderKey, SortDir};
7
8pub use toolkit_odata::ODataQuery;
10#[derive(Deserialize, Default)]
18pub struct ODataParams {
19 #[serde(rename = "$filter")]
20 pub filter: Option<String>,
21 #[serde(rename = "$orderby")]
22 pub orderby: Option<String>,
23 #[serde(rename = "$select")]
24 pub select: Option<String>,
25 #[serde(alias = "$top")]
34 pub limit: Option<u64>,
35 #[serde(alias = "$skiptoken")]
43 pub cursor: Option<String>,
44}
45
46pub const ACCEPTED_SYSTEM_QUERY_OPTIONS: [&str; 5] =
58 ["$filter", "$orderby", "$select", "$top", "$skiptoken"];
59
60const ODATA_SYSTEM_QUERY_OPTIONS: [&str; 12] = [
68 "$filter",
69 "$expand",
70 "$select",
71 "$orderby",
72 "$top",
73 "$skip",
74 "$count",
75 "$search",
76 "$format",
77 "$compute",
78 "$index",
79 "$schemaversion",
80];
81
82const UNSUPPORTED_QUERY_PARAM: &str = "UNSUPPORTED_QUERY_PARAM";
84
85pub const MAX_FILTER_LEN: usize = 8 * 1024;
86pub const MAX_NODES: usize = 2000;
87pub const MAX_ORDERBY_LEN: usize = 1024;
88pub const MAX_ORDER_FIELDS: usize = 10;
89pub const MAX_SELECT_LEN: usize = 2048;
90pub const MAX_SELECT_FIELDS: usize = 100;
91
92fn select_invalid_arg(detail: impl Into<String>, reason: &'static str) -> CanonicalError {
94 OdataError::invalid_argument()
95 .with_field_violation("$select", detail, reason)
96 .create()
97}
98
99#[allow(clippy::result_large_err)]
108pub fn parse_select(raw: &str) -> Result<Vec<String>, CanonicalError> {
109 let raw = raw.trim();
110 if raw.is_empty() {
111 return Err(select_invalid_arg(
112 "$select cannot be empty",
113 "INVALID_SELECT",
114 ));
115 }
116
117 if raw.len() > MAX_SELECT_LEN {
118 return Err(select_invalid_arg("$select too long", "INVALID_SELECT"));
119 }
120
121 let fields: Vec<String> = raw
122 .split(',')
123 .map(|f| f.trim().to_lowercase())
124 .filter(|f| !f.is_empty())
125 .collect();
126
127 if fields.is_empty() {
128 return Err(select_invalid_arg(
129 "$select must contain at least one field",
130 "INVALID_SELECT",
131 ));
132 }
133
134 if fields.len() > MAX_SELECT_FIELDS {
135 return Err(select_invalid_arg(
136 "$select contains too many fields",
137 "INVALID_SELECT",
138 ));
139 }
140
141 let mut seen = std::collections::HashSet::new();
143 for field in &fields {
144 if !seen.insert(field.clone()) {
145 return Err(select_invalid_arg(
146 format!("duplicate field in $select: {field}"),
147 "INVALID_SELECT",
148 ));
149 }
150 }
151
152 Ok(fields)
153}
154
155pub fn parse_orderby(raw: &str) -> Result<ODataOrderBy, toolkit_odata::Error> {
162 let raw = raw.trim();
163 if raw.is_empty() {
164 return Ok(ODataOrderBy::empty());
165 }
166
167 if raw.len() > MAX_ORDERBY_LEN {
168 return Err(toolkit_odata::Error::InvalidOrderByField(
169 "orderby too long".into(),
170 ));
171 }
172
173 let mut keys = Vec::new();
174
175 for part in raw.split(',') {
176 let part = part.trim();
177 if part.is_empty() {
178 continue;
179 }
180
181 let tokens: Vec<&str> = part.split_whitespace().collect();
182 let (field, dir) = match tokens.as_slice() {
183 [field] | [field, "asc"] => (*field, SortDir::Asc),
184 [field, "desc"] => (*field, SortDir::Desc),
185 _ => {
186 return Err(toolkit_odata::Error::InvalidOrderByField(format!(
187 "invalid orderby clause: {part}"
188 )));
189 }
190 };
191
192 if field.is_empty() {
193 return Err(toolkit_odata::Error::InvalidOrderByField(
194 "empty field name in orderby".into(),
195 ));
196 }
197
198 keys.push(OrderKey {
199 field: field.to_owned(),
200 dir,
201 });
202 }
203
204 if keys.len() > MAX_ORDER_FIELDS {
205 return Err(toolkit_odata::Error::InvalidOrderByField(
206 "too many order fields".into(),
207 ));
208 }
209
210 Ok(ODataOrderBy(keys))
211}
212
213fn filter_invalid_arg(detail: impl Into<String>, reason: &'static str) -> CanonicalError {
215 OdataError::invalid_argument()
216 .with_field_violation("$filter", detail, reason)
217 .create()
218}
219
220fn query_params_invalid_arg(detail: impl Into<String>) -> CanonicalError {
223 OdataError::invalid_argument()
224 .with_field_violation("query", detail, "INVALID_QUERY_PARAMS")
225 .create()
226}
227
228fn unsupported_option_hint(option: &str) -> &'static str {
232 match option {
233 "$skip" => {
234 "this platform pages by cursor, not by offset: send the previous \
235 page's `next_cursor` as `$skiptoken` (alias `cursor`)"
236 }
237 "$count" => {
238 "a page carries no total: `PageInfo` is \
239 `{next_cursor, prev_cursor, limit}`"
240 }
241 "$format" => {
242 "responses are JSON: negotiate the media type with the `Accept` \
243 header"
244 }
245 _ => "not implemented by this platform",
246 }
247}
248
249fn unsupported_option_detail(option: &str) -> String {
251 let canonical = option.to_ascii_lowercase();
252 if ACCEPTED_SYSTEM_QUERY_OPTIONS.contains(&canonical.as_str()) {
253 return format!("`{option}` binds only as `{canonical}`");
257 }
258 if ODATA_SYSTEM_QUERY_OPTIONS.contains(&canonical.as_str()) {
259 return format!(
260 "unsupported OData system query option `{option}`; {}",
261 unsupported_option_hint(&canonical)
262 );
263 }
264 format!(
265 "unknown query option `{option}`; the `$` prefix is reserved for \
266 OData system query options"
267 )
268}
269
270#[allow(clippy::result_large_err)]
286fn reject_unsupported_system_query_options(
287 pairs: &[(String, String)],
288) -> Result<(), CanonicalError> {
289 let mut offending: Vec<&str> = Vec::new();
293 for key in pairs.iter().map(|(key, _)| key.as_str()) {
294 if key.starts_with('$')
295 && !ACCEPTED_SYSTEM_QUERY_OPTIONS.contains(&key)
296 && !offending.contains(&key)
297 {
298 offending.push(key);
299 }
300 }
301
302 let mut offenders = offending.into_iter();
303 let Some(first) = offenders.next() else {
304 return Ok(());
305 };
306
307 let mut violations = OdataError::invalid_argument().with_field_violation(
308 first,
309 unsupported_option_detail(first),
310 UNSUPPORTED_QUERY_PARAM,
311 );
312 for option in offenders {
313 violations = violations.with_field_violation(
314 option,
315 unsupported_option_detail(option),
316 UNSUPPORTED_QUERY_PARAM,
317 );
318 }
319
320 Err(violations.create())
321}
322
323pub async fn extract_odata_query<S>(
335 parts: &mut Parts,
336 state: &S,
337) -> Result<ODataQuery, CanonicalError>
338where
339 S: Send + Sync,
340{
341 let Query(pairs) = Query::<Vec<(String, String)>>::from_request_parts(parts, state)
346 .await
347 .map_err(|e| query_params_invalid_arg(format!("Invalid query parameters: {e}")))?;
348 reject_unsupported_system_query_options(&pairs)?;
349
350 let Query(params) = Query::<ODataParams>::from_request_parts(parts, state)
351 .await
352 .map_err(|e| query_params_invalid_arg(format!("Invalid query parameters: {e}")))?;
353
354 let mut query = ODataQuery::new();
355
356 if let Some(raw_filter) = params.filter.as_ref() {
358 let raw = raw_filter.trim();
359 if !raw.is_empty() {
360 if raw.len() > MAX_FILTER_LEN {
361 return Err(filter_invalid_arg("Filter too long", "FILTER_TOO_LONG"));
362 }
363
364 let parsed = toolkit_odata::parse_filter_string(raw).map_err(|e| {
366 tracing::debug!(error = %e, filter_len = raw.len(), "OData filter parsing failed");
369 CanonicalError::from(e)
370 })?;
371
372 if parsed.node_count() > MAX_NODES {
373 tracing::debug!(
374 node_count = parsed.node_count(),
375 max_nodes = MAX_NODES,
376 "Filter complexity budget exceeded"
377 );
378 return Err(filter_invalid_arg(
379 "Filter too complex",
380 "FILTER_TOO_COMPLEX",
381 ));
382 }
383
384 let filter_hash = toolkit_odata::pagination::short_filter_hash(Some(parsed.as_expr()));
386
387 let core_expr = parsed.into_expr();
389
390 query = query.with_filter(core_expr);
391 if let Some(hash) = filter_hash {
392 query = query.with_filter_hash(hash);
393 }
394 }
395 }
396
397 if params.cursor.is_some() && params.orderby.is_some() {
399 return Err(ODataError::OrderWithCursor.into());
400 }
401
402 if let Some(cursor_str) = params.cursor.as_ref() {
404 let cursor = CursorV1::decode(cursor_str).map_err(|_| ODataError::InvalidCursor)?;
405 query = query.with_cursor(cursor);
406 query = query.with_order(ODataOrderBy::empty());
408 } else if let Some(raw_orderby) = params.orderby.as_ref() {
409 let order = parse_orderby(raw_orderby).map_err(CanonicalError::from)?;
411 query = query.with_order(order);
412 }
413
414 if let Some(limit) = params.limit {
416 if limit == 0 {
417 return Err(ODataError::InvalidLimit.into());
418 }
419 query = query.with_limit(limit);
420 }
421
422 if let Some(raw_select) = params.select.as_ref() {
424 let fields = parse_select(raw_select)?;
425 query = query.with_select(fields);
426 }
427
428 Ok(query)
429}
430
431use std::ops::Deref;
432
433#[derive(Debug, Clone)]
438pub struct OData(pub ODataQuery);
439
440impl OData {
441 #[inline]
442 pub fn into_inner(self) -> ODataQuery {
443 self.0
444 }
445}
446
447impl Deref for OData {
448 type Target = ODataQuery;
449 #[inline]
450 fn deref(&self) -> &Self::Target {
451 &self.0
452 }
453}
454
455impl AsRef<ODataQuery> for OData {
456 #[inline]
457 fn as_ref(&self) -> &ODataQuery {
458 &self.0
459 }
460}
461
462impl From<OData> for ODataQuery {
463 #[inline]
464 fn from(x: OData) -> Self {
465 x.0
466 }
467}
468
469impl<S> FromRequestParts<S> for OData
470where
471 S: Send + Sync,
472{
473 type Rejection = CanonicalError;
474
475 #[allow(clippy::manual_async_fn)]
476 fn from_request_parts(
477 parts: &mut Parts,
478 state: &S,
479 ) -> impl core::future::Future<Output = Result<Self, Self::Rejection>> + Send {
480 async move {
481 let query = extract_odata_query(parts, state).await?;
482 Ok(OData(query))
483 }
484 }
485}
486
487#[cfg(test)]
488#[cfg_attr(coverage_nightly, coverage(off))]
489#[path = "odata_tests.rs"]
490mod odata_tests;