1#![allow(clippy::result_large_err)]
2
3use anyhow::Result;
4use opentelemetry_semantic_conventions::attribute as otel_attribute;
5use spin_core::wasmtime::component::{Accessor, FutureReader, Resource, StreamReader};
6use spin_telemetry::traces::{self, Blame};
7use spin_world::MAX_HOST_BUFFERED_BYTES;
8use spin_world::spin::postgres3_0_0::postgres::{self as v3};
9use spin_world::spin::postgres4_2_0::postgres::{self as v4};
10use spin_world::v1::postgres as v1;
11use spin_world::v1::rdbms_types as v1_types;
12use spin_world::v2::postgres::{self as v2};
13use spin_world::v2::rdbms_types as v2_types;
14use tracing::Level;
15use tracing::field::Empty;
16use tracing::instrument;
17
18use crate::InstanceState;
19use crate::allowed_hosts::AllowedHostChecker;
20use crate::client::{Client, ClientFactory, HashableCertificate, QueryAsyncResult};
21
22impl<CF: ClientFactory> InstanceState<CF> {
23 async fn open_connection<Conn: 'static>(
24 &mut self,
25 address: &str,
26 root_ca: Option<HashableCertificate>,
27 ) -> Result<Resource<Conn>, v4::Error> {
28 let permit = self.semaphore.acquire().await.map_err(|_| {
29 let err = v4::Error::ConnectionFailed("too many connections".into());
30 traces::mark_as_error(&err, Some(Blame::Guest));
31 err
32 })?;
33 let client = self
34 .client_factory
35 .get_client(address, root_ca)
36 .await
37 .map_err(|e| {
38 let err = v4::Error::ConnectionFailed(format!("{e:?}"));
42 traces::mark_as_error(&err, Some(Blame::Guest));
43 err
44 })?;
45 self.connections
46 .push((client, permit))
47 .map_err(|_| {
48 let err = v4::Error::ConnectionFailed("too many connections".into());
50 traces::mark_as_error(&err, Some(Blame::Guest));
51 err
52 })
53 .map(Resource::new_own)
54 }
55
56 async fn get_client<Conn: 'static>(
57 &self,
58 connection: Resource<Conn>,
59 ) -> Result<&CF::Client, v4::Error> {
60 self.connections
61 .get(connection.rep())
62 .map(|(client, _permit)| client)
63 .ok_or_else(|| {
64 let err = v4::Error::ConnectionFailed("no connection found".into());
67 traces::mark_as_error(&err, Some(Blame::Host));
68 err
69 })
70 }
71
72 fn allowed_host_checker(&self) -> AllowedHostChecker {
73 self.allowed_host_checker.clone()
74 }
75
76 #[allow(clippy::result_large_err)]
77 async fn ensure_address_allowed(&self, address: &str) -> Result<(), v4::Error> {
78 self.allowed_host_checker
79 .ensure_address_allowed(address)
80 .await
81 }
82}
83
84fn v2_params_to_v3(
85 params: Vec<v2_types::ParameterValue>,
86) -> Result<Vec<v4::ParameterValue>, v2::Error> {
87 params.into_iter().map(|p| p.try_into()).collect()
88}
89
90fn v3_params_to_v4(params: Vec<v3::ParameterValue>) -> Vec<v4::ParameterValue> {
91 params.into_iter().map(|p| p.into()).collect()
92}
93
94impl<CF: ClientFactory> v3::HostConnection for InstanceState<CF> {
95 #[instrument(name = "spin_outbound_pg.open", skip(self, address), err(level = Level::INFO),
96 fields(otel.kind = "client", {otel_attribute::DB_SYSTEM_NAME} = "postgresql", {otel_attribute::SERVER_ADDRESS} = Empty, {otel_attribute::SERVER_PORT} = Empty, {otel_attribute::DB_NAMESPACE} = Empty))]
97 async fn open(&mut self, address: String) -> Result<Resource<v3::Connection>, v3::Error> {
98 spin_factor_outbound_networking::record_address_fields(&address);
99
100 self.ensure_address_allowed(&address)
101 .await
102 .map_err(v3::Error::from)
103 .map_err(track_address_check_error_v3)?;
104
105 self.open_connection(&address, None)
106 .await
107 .map_err(v3::Error::from)
108 }
109
110 #[instrument(name = "spin_outbound_pg.execute", skip(self, connection, params), err(level = Level::INFO),
111 fields(otel.kind = "client", {otel_attribute::DB_SYSTEM_NAME} = "postgresql"))]
112 async fn execute(
113 &mut self,
114 connection: Resource<v3::Connection>,
115 statement: String,
116 params: Vec<v3::ParameterValue>,
117 ) -> Result<u64, v3::Error> {
118 self.get_client(connection)
119 .await
120 .map_err(v3::Error::from)?
121 .execute(statement, v3_params_to_v4(params))
122 .await
123 .map_err(v3::Error::from)
124 .map_err(track_db_error_on_span_v3)
125 }
126
127 #[instrument(name = "spin_outbound_pg.query", skip(self, connection, params), err(level = Level::INFO),
128 fields(otel.kind = "client", {otel_attribute::DB_SYSTEM_NAME} = "postgresql"))]
129 async fn query(
130 &mut self,
131 connection: Resource<v3::Connection>,
132 statement: String,
133 params: Vec<v3::ParameterValue>,
134 ) -> Result<v3::RowSet, v3::Error> {
135 let rowset = self
136 .get_client(connection)
137 .await
138 .map_err(v3::Error::from)?
139 .query(statement, v3_params_to_v4(params), MAX_HOST_BUFFERED_BYTES)
140 .await
141 .map_err(v3::Error::from)
142 .map_err(track_db_error_on_span_v3)?;
143 Ok(rowset.into())
144 }
145
146 async fn drop(&mut self, connection: Resource<v3::Connection>) -> anyhow::Result<()> {
147 self.connections.remove(connection.rep());
148 Ok(())
149 }
150}
151
152pub(crate) struct ConnectionBuilder {
153 address: String,
154 root_ca: Option<HashableCertificate>,
155}
156
157impl<CF: ClientFactory> v4::HostConnectionBuilder for InstanceState<CF> {
158 async fn new(&mut self, address: String) -> Result<Resource<v4::ConnectionBuilder>> {
159 let builder = ConnectionBuilder {
160 address,
161 root_ca: None,
162 };
163 let rep = self
164 .builders
165 .push(builder)
166 .map_err(|_| anyhow::anyhow!("out of builder table space"))?;
167 let rsrc = Resource::new_own(rep);
168 Ok(rsrc)
169 }
170
171 async fn set_ca_root(
172 &mut self,
173 self_: Resource<v4::ConnectionBuilder>,
174 certificate: String,
175 ) -> Result<(), v4::Error> {
176 let root_ca = HashableCertificate::from_pem(&certificate).map_err(|e| {
177 let err = v4::Error::Other(format!("invalid root certificate: {e}"));
178 traces::mark_as_error(&err, Some(Blame::Guest));
179 err
180 })?;
181 let builder = self.builders.get_mut(self_.rep()).ok_or_else(|| {
182 let err = v4::Error::ConnectionFailed("no builder found".into());
183 traces::mark_as_error(&err, Some(Blame::Host));
184 err
185 })?;
186 builder.root_ca = Some(root_ca);
187 Ok(())
188 }
189
190 async fn build(
191 &mut self,
192 self_: Resource<v4::ConnectionBuilder>,
193 ) -> Result<Resource<v4::Connection>, v4::Error> {
194 let (address, root_ca) = self.get_builder_info(self_.rep())?;
195 self.open_connection(&address, root_ca).await
196 }
197
198 async fn drop(&mut self, builder: Resource<v4::ConnectionBuilder>) -> Result<()> {
199 self.builders.remove(builder.rep());
200 Ok(())
201 }
202}
203
204impl<CF: ClientFactory> v4::HostConnection for InstanceState<CF> {
205 #[instrument(name = "spin_outbound_pg.open", skip(self, address), err(level = Level::INFO),
206 fields(otel.kind = "client", {otel_attribute::DB_SYSTEM_NAME} = "postgresql", {otel_attribute::SERVER_ADDRESS} = Empty, {otel_attribute::SERVER_PORT} = Empty, {otel_attribute::DB_NAMESPACE} = Empty))]
207 async fn open(&mut self, address: String) -> Result<Resource<v4::Connection>, v4::Error> {
208 spin_factor_outbound_networking::record_address_fields(&address);
209
210 self.ensure_address_allowed(&address)
211 .await
212 .map_err(track_address_check_error_v4)?;
213
214 self.open_connection(&address, None).await
215 }
216
217 #[instrument(name = "spin_outbound_pg.execute", skip(self, connection, params), err(level = Level::INFO),
218 fields(otel.kind = "client", {otel_attribute::DB_SYSTEM_NAME} = "postgresql"))]
219 async fn execute(
220 &mut self,
221 connection: Resource<v4::Connection>,
222 statement: String,
223 params: Vec<v4::ParameterValue>,
224 ) -> Result<u64, v4::Error> {
225 self.get_client(connection)
226 .await?
227 .execute(statement, params)
228 .await
229 .map_err(track_db_error_on_span_v4)
230 }
231
232 #[instrument(name = "spin_outbound_pg.query", skip(self, connection, params), err(level = Level::INFO),
233 fields(otel.kind = "client", {otel_attribute::DB_SYSTEM_NAME} = "postgresql"))]
234 async fn query(
235 &mut self,
236 connection: Resource<v4::Connection>,
237 statement: String,
238 params: Vec<v4::ParameterValue>,
239 ) -> Result<v4::RowSet, v4::Error> {
240 self.get_client(connection)
241 .await?
242 .query(statement, params, MAX_HOST_BUFFERED_BYTES)
243 .await
244 .map_err(track_db_error_on_span_v4)
245 }
246
247 async fn drop(&mut self, connection: Resource<v4::Connection>) -> anyhow::Result<()> {
248 self.connections.remove(connection.rep());
249 Ok(())
250 }
251}
252
253impl<T, CF: ClientFactory> spin_world::spin::postgres4_2_0::postgres::HostConnectionWithStore<T>
254 for crate::PgFactorData<CF>
255{
256 #[instrument(name = "spin_outbound_pg.open_async", skip(accessor, address), err(level = Level::INFO),
257 fields(otel.kind = "client", {otel_attribute::DB_SYSTEM_NAME} = "postgresql", {otel_attribute::SERVER_ADDRESS} = Empty, {otel_attribute::SERVER_PORT} = Empty, {otel_attribute::DB_NAMESPACE} = Empty))]
258 async fn open_async(
259 accessor: &Accessor<T, Self>,
260 address: String,
261 ) -> Result<Resource<v4::Connection>, v4::Error> {
262 spin_factor_outbound_networking::record_address_fields(&address);
263
264 Self::ensure_address_allowed_async(accessor, &address)
265 .await
266 .map_err(track_address_check_error_v4)?;
267 Self::open_connection_async(accessor, &address, None).await
268 }
269
270 #[instrument(name = "spin_outbound_pg.execute", skip(accessor, connection, params), err(level = Level::INFO),
271 fields(otel.kind = "client", {otel_attribute::DB_SYSTEM_NAME} = "postgresql"))]
272 async fn execute_async(
273 accessor: &Accessor<T, Self>,
274 connection: Resource<v4::Connection>,
275 statement: String,
276 params: Vec<v4::ParameterValue>,
277 ) -> Result<u64, v4::Error> {
278 let client = accessor.with(|mut access| {
279 let host = access.get();
280 host.connections
281 .get(connection.rep())
282 .map(|(client, _permit)| client.clone())
283 .unwrap()
284 });
285
286 client
287 .execute(statement, params)
288 .await
289 .map_err(track_db_error_on_span_v4)
290 }
291
292 #[allow(clippy::type_complexity)] #[instrument(name = "spin_outbound_pg.query_async", skip(accessor, params), err(level = Level::INFO),
294 fields(otel.kind = "client", {otel_attribute::DB_SYSTEM_NAME} = "postgresql"))]
295 async fn query_async(
296 accessor: &Accessor<T, Self>,
297 connection: Resource<v4::Connection>,
298 statement: String,
299 params: Vec<v4::ParameterValue>,
300 ) -> Result<
301 (
302 Vec<v4::Column>,
303 StreamReader<v4::Row>,
304 FutureReader<Result<(), v4::Error>>,
305 ),
306 v4::Error,
307 > {
308 let client = accessor.with(|mut access| {
309 let host = access.get();
310 host.connections
311 .get(connection.rep())
312 .map(|(client, _permit)| client.clone())
313 .unwrap()
314 });
315
316 let QueryAsyncResult {
317 columns,
318 rows,
319 error,
320 } = client
321 .query_async(statement, params, MAX_HOST_BUFFERED_BYTES)
322 .await
323 .map_err(track_db_error_on_span_v4)?;
324
325 let row_producer = spin_wasi_async::stream::producer(rows);
326
327 let (sr, efr) = accessor
328 .with(|mut access| {
329 let sr = StreamReader::new(&mut access, row_producer)?;
330 let efr = FutureReader::new(&mut access, error)?;
331 anyhow::Ok((sr, efr))
332 })
333 .map_err(|e| {
334 let err = v4::Error::Other(e.to_string());
337 traces::mark_as_error(&err, Some(Blame::Host));
338 err
339 })?;
340
341 Ok((columns, sr, efr))
342 }
343}
344
345impl<CF: ClientFactory> InstanceState<CF> {
346 #[allow(clippy::result_large_err)]
347 fn get_builder_info(
348 &mut self,
349 builder_rep: u32,
350 ) -> Result<(String, Option<HashableCertificate>), v4::Error> {
351 let builder = self.builders.get_mut(builder_rep).ok_or_else(|| {
352 let err = v4::Error::ConnectionFailed("no builder found".into());
353 traces::mark_as_error(&err, Some(Blame::Host));
354 err
355 })?;
356
357 let address = builder.address.clone();
358 let root_ca = builder.root_ca.clone();
359
360 Ok((address, root_ca))
361 }
362}
363
364impl<CF: ClientFactory> crate::PgFactorData<CF> {
365 #[allow(clippy::result_large_err)]
366 fn get_builder_info<T>(
367 accessor: &Accessor<T, Self>,
368 builder: Resource<v4::ConnectionBuilder>,
369 ) -> Result<(String, Option<HashableCertificate>), v4::Error> {
370 let builder_rep = builder.rep();
371 accessor.with(|mut access| {
372 let host = access.get();
373 host.get_builder_info(builder_rep)
374 })
375 }
376
377 async fn ensure_address_allowed_async<T>(
378 accessor: &Accessor<T, Self>,
379 address: &str,
380 ) -> Result<(), v4::Error> {
381 let allowed_host_checker = accessor.with(|mut access| {
383 let host = access.get();
384 host.allowed_host_checker()
385 });
386
387 allowed_host_checker.ensure_address_allowed(address).await
388 }
389
390 async fn open_connection_async<T>(
391 accessor: &Accessor<T, Self>,
392 address: &str,
393 root_ca: Option<HashableCertificate>,
394 ) -> Result<Resource<v4::Connection>, v4::Error> {
395 let (cf, semaphore) = accessor.with(|mut access| {
396 let host = access.get();
397 (host.client_factory.clone(), host.semaphore.clone())
398 });
399
400 let permit = semaphore.acquire().await.map_err(|_| {
401 let err = v4::Error::ConnectionFailed("too many connections".into());
402 traces::mark_as_error(&err, Some(Blame::Guest));
403 err
404 })?;
405
406 let client = cf.get_client(address, root_ca).await.map_err(|e| {
407 let err = v4::Error::ConnectionFailed(format!("{e:?}"));
408 traces::mark_as_error(&err, Some(Blame::Guest));
409 err
410 })?;
411
412 accessor.with(|mut access| {
413 let host = access.get();
414 host.connections
415 .push((client, permit))
416 .map_err(|_| {
417 let err = v4::Error::ConnectionFailed("too many connections".into());
418 traces::mark_as_error(&err, Some(Blame::Guest));
419 err
420 })
421 .map(Resource::new_own)
422 })
423 }
424}
425
426impl<T, CF: ClientFactory>
427 spin_world::spin::postgres4_2_0::postgres::HostConnectionBuilderWithStore<T>
428 for crate::PgFactorData<CF>
429{
430 async fn build_async(
431 accessor: &Accessor<T, Self>,
432 builder: Resource<v4::ConnectionBuilder>,
433 ) -> Result<Resource<v4::Connection>, v4::Error> {
434 let (address, root_ca) = Self::get_builder_info(accessor, builder)?;
435
436 spin_factor_outbound_networking::record_address_fields(&address);
437
438 Self::ensure_address_allowed_async(accessor, &address)
439 .await
440 .map_err(track_address_check_error_v4)?;
441 Self::open_connection_async(accessor, &address, root_ca).await
442 }
443}
444
445impl<CF: ClientFactory> v2_types::Host for InstanceState<CF> {
446 fn convert_error(&mut self, error: v2::Error) -> Result<v2::Error> {
447 Ok(error)
448 }
449}
450
451impl<CF: ClientFactory> v3::Host for InstanceState<CF> {
452 fn convert_error(&mut self, error: v3::Error) -> Result<v3::Error> {
453 Ok(error)
454 }
455}
456
457impl<CF: ClientFactory> v4::Host for InstanceState<CF> {
458 fn convert_error(&mut self, error: v4::Error) -> Result<v4::Error> {
459 Ok(error)
460 }
461}
462
463macro_rules! delegate {
465 ($self:ident.$name:ident($address:expr, $($arg:expr),*)) => {{
466 $self.ensure_address_allowed(&$address).await?;
467 let connection = match $self.open_connection(&$address, None).await {
468 Ok(c) => c,
469 Err(e) => return Err(e.into()),
470 };
471 let rep = connection.rep();
474 let result = <Self as v4::HostConnection>::$name($self, connection, $($arg),*)
475 .await
476 .map_err(|e| e.into());
477 $self.connections.remove(rep);
478 result
479 }};
480}
481
482impl<CF: ClientFactory> v2::Host for InstanceState<CF> {}
483
484impl<CF: ClientFactory> v2::HostConnection for InstanceState<CF> {
485 #[instrument(name = "spin_outbound_pg.open", skip(self, address), err(level = Level::INFO),
486 fields(otel.kind = "client", {otel_attribute::DB_SYSTEM_NAME} = "postgresql", {otel_attribute::SERVER_ADDRESS} = Empty, {otel_attribute::SERVER_PORT} = Empty, {otel_attribute::DB_NAMESPACE} = Empty))]
487 async fn open(&mut self, address: String) -> Result<Resource<v2::Connection>, v2::Error> {
488 self.otel.reparent_tracing_span();
489 spin_factor_outbound_networking::record_address_fields(&address);
490
491 self.ensure_address_allowed(&address)
492 .await
493 .map_err(v2::Error::from)
494 .map_err(track_address_check_error_v2)?;
495 self.open_connection(&address, None)
496 .await
497 .map_err(v2::Error::from)
498 }
499
500 #[instrument(name = "spin_outbound_pg.execute", skip(self, connection, params), err(level = Level::INFO),
501 fields(otel.kind = "client", {otel_attribute::DB_SYSTEM_NAME} = "postgresql"))]
502 async fn execute(
503 &mut self,
504 connection: Resource<v2::Connection>,
505 statement: String,
506 params: Vec<v2_types::ParameterValue>,
507 ) -> Result<u64, v2::Error> {
508 self.otel.reparent_tracing_span();
509 let params = v2_params_to_v3(params).inspect_err(|e| {
510 traces::mark_as_error(e, Some(Blame::Guest));
511 })?;
512 self.get_client(connection)
513 .await
514 .map_err(v2::Error::from)?
515 .execute(statement, params)
516 .await
517 .map_err(v2::Error::from)
518 .map_err(track_db_error_on_span_v2)
519 }
520
521 #[instrument(name = "spin_outbound_pg.query", skip(self, connection, params), err(level = Level::INFO),
522 fields(otel.kind = "client", {otel_attribute::DB_SYSTEM_NAME} = "postgresql"))]
523 async fn query(
524 &mut self,
525 connection: Resource<v2::Connection>,
526 statement: String,
527 params: Vec<v2_types::ParameterValue>,
528 ) -> Result<v2_types::RowSet, v2::Error> {
529 self.otel.reparent_tracing_span();
530 let params = v2_params_to_v3(params).inspect_err(|e| {
531 traces::mark_as_error(e, Some(Blame::Guest));
532 })?;
533 Ok(self
534 .get_client(connection)
535 .await
536 .map_err(v2::Error::from)?
537 .query(statement, params, MAX_HOST_BUFFERED_BYTES)
538 .await
539 .map_err(v2::Error::from)
540 .map_err(track_db_error_on_span_v2)?
541 .into())
542 }
543
544 async fn drop(&mut self, connection: Resource<v2::Connection>) -> anyhow::Result<()> {
545 self.connections.remove(connection.rep());
546 Ok(())
547 }
548}
549
550impl<CF: ClientFactory> v1::Host for InstanceState<CF> {
551 async fn execute(
552 &mut self,
553 address: String,
554 statement: String,
555 params: Vec<v1_types::ParameterValue>,
556 ) -> Result<u64, v1::PgError> {
557 delegate!(
558 self.execute(
559 address,
560 statement,
561 params
562 .into_iter()
563 .map(TryInto::try_into)
564 .collect::<Result<Vec<_>, _>>()?
565 )
566 )
567 }
568
569 async fn query(
570 &mut self,
571 address: String,
572 statement: String,
573 params: Vec<v1_types::ParameterValue>,
574 ) -> Result<v1_types::RowSet, v1::PgError> {
575 delegate!(
576 self.query(
577 address,
578 statement,
579 params
580 .into_iter()
581 .map(TryInto::try_into)
582 .collect::<Result<Vec<_>, _>>()?
583 )
584 )
585 .map(Into::into)
586 }
587
588 fn convert_pg_error(&mut self, error: v1::PgError) -> Result<v1::PgError> {
589 Ok(error)
590 }
591}
592
593fn track_address_check_error_v4(err: v4::Error) -> v4::Error {
599 let blame = match &err {
600 v4::Error::Other(_) => Blame::Host,
601 _ => Blame::Guest,
602 };
603 traces::mark_as_error(&err, Some(blame));
604 err
605}
606
607fn track_address_check_error_v3(err: v3::Error) -> v3::Error {
608 let blame = match &err {
609 v3::Error::Other(_) => Blame::Host,
610 _ => Blame::Guest,
611 };
612 traces::mark_as_error(&err, Some(blame));
613 err
614}
615
616fn track_address_check_error_v2(err: v2::Error) -> v2::Error {
617 let blame = match &err {
618 v2::Error::Other(_) => Blame::Host,
619 _ => Blame::Guest,
620 };
621 traces::mark_as_error(&err, Some(blame));
622 err
623}
624
625fn track_db_error_on_span_v4(err: v4::Error) -> v4::Error {
627 let blame = match &err {
628 v4::Error::ConnectionFailed(_) => Blame::Guest,
632 v4::Error::BadParameter(_) => Blame::Guest,
633 v4::Error::QueryFailed(_) => Blame::Guest,
634 v4::Error::ValueConversionFailed(_) => Blame::Host,
637 v4::Error::Other(_) => Blame::Host,
638 };
639 traces::mark_as_error(&err, Some(blame));
640 err
641}
642
643fn track_db_error_on_span_v3(err: v3::Error) -> v3::Error {
644 let blame = match &err {
645 v3::Error::ConnectionFailed(_) => Blame::Guest,
646 v3::Error::BadParameter(_) => Blame::Guest,
647 v3::Error::QueryFailed(_) => Blame::Guest,
648 v3::Error::ValueConversionFailed(_) => Blame::Host,
649 v3::Error::Other(_) => Blame::Host,
650 };
651 traces::mark_as_error(&err, Some(blame));
652 err
653}
654
655fn track_db_error_on_span_v2(err: v2::Error) -> v2::Error {
656 let blame = match &err {
657 v2::Error::ConnectionFailed(_) => Blame::Guest,
658 v2::Error::BadParameter(_) => Blame::Guest,
659 v2::Error::QueryFailed(_) => Blame::Guest,
660 v2::Error::ValueConversionFailed(_) => Blame::Host,
661 v2::Error::Other(_) => Blame::Host,
662 };
663 traces::mark_as_error(&err, Some(blame));
664 err
665}