Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
27 changes: 25 additions & 2 deletions pgdog/src/frontend/client/query_engine/advisory_lock.rs
Original file line number Diff line number Diff line change
@@ -1,17 +1,39 @@
use fnv::FnvHashSet;

use crate::frontend::router::parser::statement::{
AdvisoryLockId, AdvisoryLocks as ParserAdvisoryLocks, LockScope,
use crate::{
frontend::router::parser::statement::{
AdvisoryLockId, AdvisoryLocks as ParserAdvisoryLocks, LockScope,
},
net::{DataRow, Error, FromBytes, Message, ToBytes},
};

/// Tracks advisory locks held by the current client across requests.
#[derive(Default, Debug)]
pub(crate) struct AdvisoryLocks {
locks: FnvHashSet<AdvisoryLockId>,
/// pg_try_advisory_lock returned false.
not_acquired: bool,
}

impl AdvisoryLocks {
/// Check if pg_try_advisory_lock acquired the lock.
pub(crate) fn data_row(
&mut self,
locks: &ParserAdvisoryLocks,
message: &Message,
) -> Result<(), Error> {
if locks.try_lock() {
let row = DataRow::from_bytes(message.to_bytes())?;
// Text format is 't' / 'f', binary is 1 / 0.
self.not_acquired = !matches!(row.column(0).as_deref(), Some(b"t" | b"\x01"));
}

Ok(())
}

pub(crate) fn merge(&mut self, locks: &ParserAdvisoryLocks) {
let not_acquired = std::mem::take(&mut self.not_acquired) && locks.try_lock();

for lock in locks.iter() {
if lock.unlock_all {
self.locks.clear();
Expand All @@ -22,6 +44,7 @@ impl AdvisoryLocks {
}
} else if let Some(id) = lock.id
&& lock.scope == LockScope::Session
&& !not_acquired
{
self.locks.insert(id);
}
Expand Down
5 changes: 5 additions & 0 deletions pgdog/src/frontend/client/query_engine/query.rs
Original file line number Diff line number Diff line change
Expand Up @@ -193,6 +193,11 @@ impl QueryEngine {
self.emit_explain_rows(context).await?;
}

if code == 'D' {
self.advisory_locks
.data_row(self.router.command().route().advisory_locks(), &message)?;
}

if code == 'E' {
if let Some(state) = self.pending_explain.as_mut() {
state.annotated = true;
Expand Down
86 changes: 85 additions & 1 deletion pgdog/src/frontend/client/query_engine/test/advisory_lock.rs
Original file line number Diff line number Diff line change
@@ -1,5 +1,9 @@
use super::prelude::*;
use crate::{frontend::router::parser::statement::AdvisoryLockId, net::DataRow};
use crate::{
expect_message,
frontend::router::parser::statement::AdvisoryLockId,
net::{BindComplete, CommandComplete, DataRow, Format, ParseComplete, ReadyForQuery},
};

#[tokio::test]
async fn test_pg_catalog_advisory_lock_pins_until_qualified_unlock() {
Expand Down Expand Up @@ -432,3 +436,83 @@ async fn test_xact_lock_released_on_rollback() {
assert_eq!(locks.len(), 0);
assert!(!client.backend_locked());
}

#[tokio::test]
async fn test_failed_try_lock_does_not_pin() {
let mut holder = TestClient::new_sharded(Parameters::default()).await;
let mut client = TestClient::new(Parameters::default()).await;

holder
.send_simple(Query::new("SELECT pg_advisory_lock(2026100501)"))
.await;
holder.read_until('Z').await.unwrap();

client
.send_simple(Query::new("SELECT pg_try_advisory_lock(2026100501)"))
.await;
let messages = client.read_until('Z').await.unwrap();
let row = messages
.iter()
.find(|m| m.code() == 'D')
.map(|m| DataRow::try_from(m.clone()).unwrap())
.unwrap();
assert_eq!(row.get_text(0).as_deref(), Some("f"));

assert_eq!(client.engine.advisory_locks().len(), 0);
assert!(!client.backend_locked());

// Lock is free now, so it should pin.
holder
.send_simple(Query::new("SELECT pg_advisory_unlock(2026100501)"))
.await;
holder.read_until('Z').await.unwrap();

client
.send_simple(Query::new("SELECT pg_try_advisory_lock(2026100501)"))
.await;
client.read_until('Z').await.unwrap();

assert!(
client
.engine
.advisory_locks()
.contains(AdvisoryLockId::OneParameter(2026100501))
);
assert!(client.backend_locked());
}

#[tokio::test]
async fn test_failed_try_lock_binary_does_not_pin() {
let mut holder = TestClient::new_sharded(Parameters::default()).await;
let mut client = TestClient::new(Parameters::default()).await;

holder
.send_simple(Query::new("SELECT pg_advisory_lock(2026100502)"))
.await;
holder.read_until('Z').await.unwrap();

client
.send(Parse::named("try_lock", "SELECT pg_try_advisory_lock($1)"))
.await;
client
.send(Bind::new_params_codes_results(
"try_lock",
&[Parameter::new(&2026100502_i64.to_be_bytes())],
&[Format::Binary],
&[1],
))
.await;
client.send(Execute::new()).await;
client.send(Sync).await;
client.try_process().await.unwrap();

expect_message!(client.read().await, ParseComplete);
expect_message!(client.read().await, BindComplete);
let row = expect_message!(client.read().await, DataRow);
assert_eq!(row.get::<bool>(0, Format::Binary), Some(false));
expect_message!(client.read().await, CommandComplete);
expect_message!(client.read().await, ReadyForQuery);

assert_eq!(client.engine.advisory_locks().len(), 0);
assert!(!client.backend_locked());
}
76 changes: 73 additions & 3 deletions pgdog/src/frontend/router/parser/statement.rs
Original file line number Diff line number Diff line change
Expand Up @@ -244,7 +244,7 @@ fn is_param_ref(node: Node<'_>) -> bool {
}

use super::{
super::sharding::Value as ShardingValue, Column, Error, Table, Value,
super::sharding::Value as ShardingValue, Column, Error, Function, Table, Value,
explain_trace::ExplainEntry,
};

Expand Down Expand Up @@ -296,6 +296,9 @@ impl AdvisoryLockId {
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub(crate) struct AdvisoryLocks {
locks: HashSet<AdvisoryLock>,
/// The statement is `SELECT pg_try_advisory_lock(...)` and nothing else,
/// so the first column tells us if the lock was acquired.
try_lock: bool,
}

impl AdvisoryLocks {
Expand All @@ -306,6 +309,10 @@ impl AdvisoryLocks {
pub(crate) fn is_empty(&self) -> bool {
self.locks.is_empty()
}

pub(crate) fn try_lock(&self) -> bool {
self.try_lock
}
}

/// Accumulator shared across statement walkers — lets a single traversal
Expand Down Expand Up @@ -758,9 +765,41 @@ impl<'a, 'b: 'a> StatementParser<'a, 'b> {

/// Extract pg_advisory_lock / pg_advisory_unlock calls with literal integer keys.
pub(crate) fn extract_advisory_locks(&mut self) -> AdvisoryLocks {
AdvisoryLocks {
locks: self.walk().advisory_locks.clone(),
let locks = self.walk().advisory_locks.clone();
let try_lock = locks.len() == 1 && self.is_try_lock();

AdvisoryLocks { locks, try_lock }
}

/// `SELECT pg_try_advisory_lock(...)` with no casts, FROM or other columns.
/// Postgres returns false instead of an error if the lock is taken.
fn is_try_lock(&self) -> bool {
let Node::SelectStmt(stmt) = self.stmt else {
return false;
};

if !stmt.from_clause().is_empty() {
return false;
}

let Ok(target) = stmt.target_list().into_iter().exactly_one() else {
return false;
};

let Node::FuncCall(func) = target.val() else {
return false;
};

let Some(func) = Function::from_strings(func.funcname().iter().filter_map(Node::as_str))
else {
return false;
};

matches!(func.schema, None | Some("pg_catalog"))
&& matches!(
func.name,
"pg_try_advisory_lock" | "pg_try_advisory_lock_shared"
)
}

// Are we running? Or walking? MAKE UP YOUR MIND DAMMIT
Expand Down Expand Up @@ -3553,5 +3592,36 @@ mod test {
],
);
}

#[test]
fn try_lock_only_for_single_column_select() {
let try_lock = |query: &str| {
let schema = ShardingSchema::default();
let raw = pg_raw_parse::parse(query).unwrap();
let stmt = raw.stmts().next().unwrap();
StatementParser::new(stmt, None, &schema)
.extract_advisory_locks()
.try_lock()
};

for query in [
"SELECT pg_try_advisory_lock(1)",
"SELECT pg_try_advisory_lock_shared(1)",
"SELECT pg_catalog.pg_try_advisory_lock(1, 2)",
] {
assert!(try_lock(query), "{query}");
}

for query in [
"SELECT pg_advisory_lock(1)",
"SELECT pg_try_advisory_xact_lock(1)",
"SELECT pg_try_advisory_lock(1)::int",
"SELECT 1, pg_try_advisory_lock(1)",
"SELECT pg_try_advisory_lock(1), pg_try_advisory_lock(2)",
"SELECT pg_try_advisory_lock(v) FROM (VALUES (1)) AS t(v)",
] {
assert!(!try_lock(query), "{query}");
}
}
}
}