use std::fmt::Debug;
use std::sync::Arc;
use async_trait::async_trait;
use futures::SinkExt;
use futures::sink::Sink;
use pgwire::api::portal::Portal;
use pgwire::api::query::ExtendedQueryHandler;
use pgwire::api::results::{DescribePortalResponse, DescribeStatementResponse, Response};
use pgwire::api::stmt::StoredStatement;
use pgwire::api::store::PortalStore;
use pgwire::api::{ClientInfo, ClientPortalStore};
use pgwire::error::{PgWireError, PgWireResult};
use pgwire::messages::PgWireBackendMessage;
use crate::config::auth::AuthMode;
use crate::control::planner::context::QueryContext;
use crate::control::security::identity::AuthenticatedIdentity;
use crate::control::state::SharedState;
use crate::types::RequestId;
use super::super::types::notice_warning;
use super::in_flight::InFlightGuard;
use super::prepared::{NodeDbQueryParser, ParsedStatement};
use crate::control::server::shared::session::SessionStore;
mod simple_query;
pub struct NodeDbPgHandler {
pub(crate) state: Arc<SharedState>,
pub(super) query_ctx: QueryContext,
query_parser: Arc<NodeDbQueryParser>,
pub(super) auth_mode: AuthMode,
pub(crate) sessions: Arc<SessionStore>,
pub(crate) restore_state: Arc<crate::control::backup::RestoreState>,
}
impl NodeDbPgHandler {
pub fn new(state: Arc<SharedState>, auth_mode: AuthMode) -> Self {
let query_ctx = QueryContext::for_state_with_lease(&state);
let sessions = Arc::new(SessionStore::new());
let query_parser = Arc::new(NodeDbQueryParser::new(
Arc::clone(&state),
auth_mode.clone(),
Arc::clone(&sessions),
));
Self {
state,
query_ctx,
query_parser,
auth_mode,
sessions,
restore_state: Arc::new(crate::control::backup::RestoreState::new()),
}
}
pub(super) fn next_request_id(&self) -> RequestId {
self.state.next_request_id()
}
pub(crate) fn resolve_identity<C: ClientInfo>(
&self,
client: &C,
addr: &std::net::SocketAddr,
) -> PgWireResult<AuthenticatedIdentity> {
super::auth::resolve_session_identity(
&self.state,
self.auth_mode.clone(),
&self.sessions,
client,
addr,
)
}
}
#[async_trait]
impl ExtendedQueryHandler for NodeDbPgHandler {
type Statement = ParsedStatement;
type QueryParser = NodeDbQueryParser;
fn query_parser(&self) -> Arc<Self::QueryParser> {
self.query_parser.clone()
}
async fn do_query<C>(
&self,
client: &mut C,
portal: &Portal<Self::Statement>,
max_rows: usize,
) -> PgWireResult<Response>
where
C: ClientInfo + ClientPortalStore + Sink<PgWireBackendMessage> + Unpin + Send + Sync,
C::PortalStore: PortalStore<Statement = Self::Statement>,
C::Error: Debug,
PgWireError: From<<C as Sink<PgWireBackendMessage>>::Error>,
{
let addr = client.socket_addr();
let _in_flight = InFlightGuard::new(&self.sessions, addr);
let result = self.execute_prepared(client, portal, max_rows).await;
for message in self.sessions.drain_notices(&addr) {
let notice = notice_warning(&message);
let _ = client
.send(PgWireBackendMessage::NoticeResponse(notice))
.await;
}
result
}
async fn do_describe_statement<C>(
&self,
client: &mut C,
target: &StoredStatement<Self::Statement>,
) -> PgWireResult<DescribeStatementResponse>
where
C: ClientInfo + ClientPortalStore + Sink<PgWireBackendMessage> + Unpin + Send + Sync,
C::PortalStore: PortalStore<Statement = Self::Statement>,
C::Error: Debug,
PgWireError: From<<C as Sink<PgWireBackendMessage>>::Error>,
{
self.describe_statement_impl(client, target).await
}
async fn do_describe_portal<C>(
&self,
client: &mut C,
target: &Portal<Self::Statement>,
) -> PgWireResult<DescribePortalResponse>
where
C: ClientInfo + ClientPortalStore + Sink<PgWireBackendMessage> + Unpin + Send + Sync,
C::PortalStore: PortalStore<Statement = Self::Statement>,
C::Error: Debug,
PgWireError: From<<C as Sink<PgWireBackendMessage>>::Error>,
{
self.describe_portal_impl(client, target).await
}
}