postrust_core/plan/
call_plan.rs1use crate::api_request::{ApiRequest, Payload, QualifiedIdentifier};
4use crate::error::{Error, Result};
5use crate::schema_cache::Routine;
6use serde::{Deserialize, Serialize};
7
8#[derive(Clone, Debug, Serialize, Deserialize)]
10pub struct CallPlan {
11 pub function: QualifiedIdentifier,
13 pub params: CallParams,
15 pub returns_scalar: bool,
17 pub returns_set: bool,
19 #[serde(default)]
22 pub returns_composite: bool,
23 pub volatility: String,
25}
26
27#[derive(Clone, Debug, Serialize, Deserialize)]
29pub enum CallParams {
30 Named(Vec<(String, String)>),
32 Positional(Vec<String>),
34 SingleObject(bytes::Bytes),
36 None,
38}
39
40impl CallPlan {
41 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 pub fn has_params(&self) -> bool {
66 !matches!(self.params, CallParams::None)
67 }
68}
69
70fn extract_call_params(request: &ApiRequest, _routine: &Routine) -> Result<CallParams> {
72 if let Some(payload) = &request.payload {
74 match payload {
75 Payload::ProcessedJson { raw, .. } => {
76 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 let params: Vec<(String, String)> = map
84 .into_iter()
85 .map(|(k, v)| {
86 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 return Ok(CallParams::SingleObject(raw.clone()));
100 }
101 _ => {
102 return Ok(CallParams::SingleObject(raw.clone()));
104 }
105 }
106 }
107 Payload::ProcessedUrlEncoded { data, .. } => {
108 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 if !request.query_params.params.is_empty() {
119 return Ok(CallParams::Named(request.query_params.params.clone()));
120 }
121
122 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}