Skip to main content

postrust_core/plan/
call_plan.rs

1//! RPC (stored function) call planning.
2
3use crate::api_request::{ApiRequest, Payload, QualifiedIdentifier};
4use crate::error::{Error, Result};
5use crate::schema_cache::Routine;
6use serde::{Deserialize, Serialize};
7
8/// A plan for calling a stored function.
9#[derive(Clone, Debug, Serialize, Deserialize)]
10pub struct CallPlan {
11    /// Function identifier
12    pub function: QualifiedIdentifier,
13    /// Call parameters
14    pub params: CallParams,
15    /// Whether to return a scalar result
16    pub returns_scalar: bool,
17    /// Whether the function is set-returning
18    pub returns_set: bool,
19    /// Whether the return type is composite (row type or `record`): its
20    /// columns are real output columns, never the function-name wrapper.
21    #[serde(default)]
22    pub returns_composite: bool,
23    /// Function volatility (for transaction handling)
24    pub volatility: String,
25}
26
27/// How parameters are passed to the function.
28#[derive(Clone, Debug, Serialize, Deserialize)]
29pub enum CallParams {
30    /// Named parameters from URL query or JSON body
31    Named(Vec<(String, String)>),
32    /// Positional parameters (from JSON array)
33    Positional(Vec<String>),
34    /// Single JSON object passed as first argument
35    SingleObject(bytes::Bytes),
36    /// No parameters
37    None,
38}
39
40impl CallPlan {
41    /// Create a call plan from an API request.
42    pub fn from_request(request: &ApiRequest, routine: &Routine) -> Result<Self> {
43        let qi = routine.qualified_identifier();
44
45        let params = extract_call_params(request, routine)?;
46
47        let returns_scalar = !routine.return_type.is_set_returning()
48            && routine
49                .return_type
50                .type_name()
51                .map(|t| !t.contains("record"))
52                .unwrap_or(true);
53
54        Ok(Self {
55            function: qi,
56            params,
57            returns_scalar,
58            returns_set: routine.return_type.is_set_returning(),
59            returns_composite: routine.returns_composite,
60            volatility: format!("{:?}", routine.volatility),
61        })
62    }
63
64    /// Check if this call has parameters.
65    pub fn has_params(&self) -> bool {
66        !matches!(self.params, CallParams::None)
67    }
68}
69
70/// Extract call parameters from request.
71fn extract_call_params(request: &ApiRequest, _routine: &Routine) -> Result<CallParams> {
72    // Check for JSON body first
73    if let Some(payload) = &request.payload {
74        match payload {
75            Payload::ProcessedJson { raw, .. } => {
76                // Check if it's an object or array
77                let value: serde_json::Value =
78                    serde_json::from_slice(raw).map_err(|e| Error::InvalidBody(e.to_string()))?;
79
80                match value {
81                    serde_json::Value::Object(map) => {
82                        // Named parameters from JSON object
83                        let params: Vec<(String, String)> = map
84                            .into_iter()
85                            .map(|(k, v)| {
86                                // Extract string values without JSON quotes
87                                let value = match v {
88                                    serde_json::Value::String(s) => s,
89                                    serde_json::Value::Null => String::new(),
90                                    other => other.to_string(),
91                                };
92                                (k, value)
93                            })
94                            .collect();
95                        return Ok(CallParams::Named(params));
96                    }
97                    serde_json::Value::Array(_) => {
98                        // Pass entire JSON as single argument
99                        return Ok(CallParams::SingleObject(raw.clone()));
100                    }
101                    _ => {
102                        // Scalar value - pass as single argument
103                        return Ok(CallParams::SingleObject(raw.clone()));
104                    }
105                }
106            }
107            Payload::ProcessedUrlEncoded { data, .. } => {
108                // Named parameters from form data
109                return Ok(CallParams::Named(data.clone()));
110            }
111            Payload::RawJson(raw) | Payload::RawPayload(raw) => {
112                return Ok(CallParams::SingleObject(raw.clone()));
113            }
114        }
115    }
116
117    // Fall back to query parameters
118    if !request.query_params.params.is_empty() {
119        return Ok(CallParams::Named(request.query_params.params.clone()));
120    }
121
122    // No parameters
123    Ok(CallParams::None)
124}
125
126#[cfg(test)]
127mod tests {
128    use super::*;
129    use crate::schema_cache::{FuncVolatility, RetType};
130
131    fn make_routine() -> Routine {
132        Routine {
133            schema: "public".into(),
134            name: "get_users".into(),
135            description: None,
136            params: vec![],
137            return_type: RetType::SetOf("users".into()),
138            returns_composite: true,
139            volatility: FuncVolatility::Stable,
140            has_variadic: false,
141            isolation_level: None,
142            settings: vec![],
143            is_procedure: false,
144        }
145    }
146
147    #[test]
148    fn test_call_plan_basic() {
149        let request = ApiRequest::default();
150        let routine = make_routine();
151
152        let plan = CallPlan::from_request(&request, &routine).unwrap();
153
154        assert_eq!(plan.function.name, "get_users");
155        assert!(plan.returns_set);
156        assert!(!plan.returns_scalar);
157        assert!(plan.returns_composite);
158    }
159
160    #[test]
161    fn test_call_plan_scalar_is_not_composite() {
162        let request = ApiRequest::default();
163        let routine = Routine {
164            return_type: RetType::Single("integer".into()),
165            returns_composite: false,
166            ..make_routine()
167        };
168
169        let plan = CallPlan::from_request(&request, &routine).unwrap();
170
171        assert!(plan.returns_scalar);
172        assert!(!plan.returns_set);
173        assert!(!plan.returns_composite);
174    }
175
176    #[test]
177    fn test_call_params_none() {
178        let request = ApiRequest::default();
179        let routine = make_routine();
180
181        let plan = CallPlan::from_request(&request, &routine).unwrap();
182        assert!(!plan.has_params());
183    }
184}