keep postgress connections open and switch tasks between them

This commit is contained in:
Bill
2025-11-18 13:08:32 -07:00
parent ac61f58f01
commit 2888c0b5e7
2 changed files with 56 additions and 61 deletions

View File

@@ -16,7 +16,7 @@ DB_PORT = "5433" # <-- Your local Docker port
DB_NAME = "postgres" DB_NAME = "postgres"
# Global Configuration # Global Configuration
BLAST_RADIUS = 312500 # Rows per parallel task (1M / 62500 = 16 partitions) BLAST_RADIUS = 125000 # Rows per parallel task (1M / 62500 = 16 partitions)
# SQLAlchemy connection string (for Pandas) # SQLAlchemy connection string (for Pandas)
sqlalchemy_conn_str = f"postgresql://{DB_USER}:{DB_PASS}@{DB_HOST}:{DB_PORT}/{DB_NAME}" sqlalchemy_conn_str = f"postgresql://{DB_USER}:{DB_PASS}@{DB_HOST}:{DB_PORT}/{DB_NAME}"

View File

@@ -1,7 +1,9 @@
// --- External Crates --- // --- External Crates ---
use std::pin::Pin; use std::pin::Pin;
use std::sync::Arc; use std::sync::Arc;
// FIX: Replace direct tokio_postgres connections with deadpool
use tokio_postgres::{NoTls, CopyOutStream, Config as PgConfig}; use tokio_postgres::{NoTls, CopyOutStream, Config as PgConfig};
use deadpool_postgres::{Pool, Manager, Runtime};
use anyhow::{Context, Result, anyhow}; use anyhow::{Context, Result, anyhow};
use arrow::array::{ use arrow::array::{
ArrayBuilder, ArrayRef, ArrayBuilder, ArrayRef,
@@ -28,40 +30,48 @@ use crate::config::ConnectorConfig;
// --- CONSTANTS --- // --- CONSTANTS ---
const POSTGRES_EPOCH_NAIVE: NaiveDateTime = NaiveDate::from_ymd_opt(2000, 1, 1).unwrap().and_hms_opt(0, 0, 0).unwrap(); // Optimized calculation of epoch delta (2000-01-01 00:00:00 to 1970-01-01 00:00:00)
const UNIX_EPOCH_NAIVE_DATE: NaiveDate = NaiveDate::from_ymd_opt(1970, 1, 1).unwrap(); // 10957 days * 86400 seconds/day * 1,000,000 micros/second = 946684800000000 micros
const POSTGRES_EPOCH_MICROS_OFFSET: i64 = 946684800000000;
// --- 1. CORE DATABASE LOGIC (PARALLEL COORDINATOR) --- // --- 1. CORE DATABASE LOGIC (PARALLEL COORDINATOR) ---
// Note: We remove the unused 'danger_mode' from the function signature to simplify the fast path API. // Note: We use blast_radius from Python config as the partitioning strategy
pub async fn run_db_logic(config: ConnectorConfig, blast_radius: i64) -> Result<RecordBatch> { pub async fn run_db_logic(config: ConnectorConfig, blast_radius: i64) -> Result<RecordBatch> {
// Removed danger_mode print, now defaults to fast path
println!("UncheckedIO: Starting Query Planner (Blast Radius: {})...", blast_radius); println!("UncheckedIO: Starting Query Planner (Blast Radius: {})...", blast_radius);
// 1. Establish the *coordinator* connection // --- PHASE 1: SETUP CONNECTION POOL ---
let pg_config = PgConfig::from_str(&config.connection_string)?; let pg_config: tokio_postgres::Config = PgConfig::from_str(&config.connection_string)
let (client, connection) = pg_config.connect(NoTls).await .context("Invalid connection string in config")?;
.context("Coordinator: Failed to connect to PostgreSQL")?;
tokio::spawn(async move {
if let Err(e) = connection.await { eprintln!("Coordinator connection error: {}", e); }
});
// 2. Define Partition Strategy // Initialize the Manager and Pool
let manager = Manager::new(pg_config.clone(), NoTls);
// Set pool size higher than the expected partition count to ensure connections are always available.
let pool = Pool::builder(manager)
.max_size(20)
.runtime(Runtime::Tokio1)
.build()
.context("Failed to build connection pool")?;
// 2. Query for Table Bounds (using pool connection)
let client = pool.get().await.context("Failed to get pool connection for stats query")?;
let partition_key = "id"; let partition_key = "id";
// 3. Query for Table Bounds
let (base_query, _) = config.query.trim().split_once("TO STDOUT (FORMAT binary)") let (base_query, _) = config.query.trim().split_once("TO STDOUT (FORMAT binary)")
.context("Failed to parse base query from config")?; .context("Failed to parse base query from config")?;
let base_query_inner = base_query.trim().trim_start_matches("COPY (").trim_end_matches(")"); let base_query_inner = base_query.trim().trim_start_matches("COPY (").trim_end_matches(")");
// NOTE: We rely on MIN/MAX here, assuming dense key for benchmark data.
let stats_query = format!("SELECT MIN({}), MAX({}) FROM ({}) AS subquery", partition_key, partition_key, base_query_inner); let stats_query = format!("SELECT MIN({}), MAX({}) FROM ({}) AS subquery", partition_key, partition_key, base_query_inner);
let row = client.query_one(&stats_query, &[]).await?; let row = client.query_one(&stats_query, &[]).await?;
let min_id: i64 = row.try_get(0).context("Failed to get MIN(id)")?; let min_id: i64 = row.try_get(0).context("Failed to get MIN(id)")?;
let max_id: i64 = row.try_get(1).context("Failed to get MAX(id)")?; let max_id: i64 = row.try_get(1).context("Failed to get MAX(id)")?;
// Connection is returned to the pool when 'client' is dropped here.
println!("UncheckedIO: ID Range: {} to {}", min_id, max_id); println!("UncheckedIO: ID Range: {} to {}", min_id, max_id);
// 4. Generate Partitioned Queries (Unchanged) // 4. Generate Partitioned Queries (Based on blast_radius from Python)
struct PartitionTask { struct PartitionTask {
index: usize, index: usize,
query: String, query: String,
@@ -91,20 +101,19 @@ pub async fn run_db_logic(config: ConnectorConfig, blast_radius: i64) -> Result<
let mut join_set = JoinSet::new(); let mut join_set = JoinSet::new();
for task in partitions { for task in partitions {
let worker_pg_config = pg_config.clone(); let worker_pool = pool.clone(); // Pass the pool handle
let worker_schema = arrow_schema.clone(); let worker_schema = arrow_schema.clone();
join_set.spawn(async move { join_set.spawn(async move {
let worker_logic = async { let worker_logic = async {
let (worker_client, worker_connection) = worker_pg_config.connect(NoTls).await?; // Get connection from pool (This is the speedup)
tokio::spawn(async move { let worker_client = worker_pool.get().await
if let Err(e) = worker_connection.await { eprintln!("Worker connection error: {}", e); } .context("Worker: Failed to get pool connection")?;
});
let copy_stream = worker_client.copy_out(task.query.as_str()).await?; let copy_stream = worker_client.copy_out(task.query.as_str()).await?;
let pinned_stream: Pin<Box<CopyOutStream>> = Box::pin(copy_stream); let pinned_stream: Pin<Box<CopyOutStream>> = Box::pin(copy_stream);
// Now call the new parser entry point // Call the static dispatch parser
parse_data_with_schema(pinned_stream, worker_schema).await parse_data_with_schema(pinned_stream, worker_schema).await
}; };
@@ -153,67 +162,45 @@ fn create_null_batch(schema: Arc<Schema>, num_rows: usize) -> Result<RecordBatch
RecordBatch::try_new(schema, columns).context("Failed to create null placeholder batch") RecordBatch::try_new(schema, columns).context("Failed to create null placeholder batch")
} }
/// Helper function to build the Arrow Schema from the config
fn build_arrow_schema(config: &ConnectorConfig) -> Result<Schema> { fn build_arrow_schema(config: &ConnectorConfig) -> Result<Schema> {
let schema_fields: Vec<Field> = config.schema.iter().map(|col_cfg| { let schema_fields: Vec<Field> = config.schema.iter().map(|col_cfg| {
let nullable = col_cfg.column_name == "notes"; let nullable = col_cfg.column_name == "notes"; // Hack for MVP
let arrow_type = match col_cfg.arrow_type.as_str() { let arrow_type = match col_cfg.arrow_type.as_str() {
"Int64" => DataType::Int64, "Int32" => DataType::Int32, "Int64" => DataType::Int64,
"Float64" => DataType::Float64, "Float32" => DataType::Float32, "Int32" => DataType::Int32,
"Utf8" | "String" => DataType::Utf8, "Boolean" => DataType::Boolean, "Float64" => DataType::Float64,
"Float32" => DataType::Float32,
"Utf8" | "String" => DataType::Utf8,
"Boolean" => DataType::Boolean,
"Timestamp(Nanosecond, None)" => DataType::Timestamp(arrow::datatypes::TimeUnit::Nanosecond, None), "Timestamp(Nanosecond, None)" => DataType::Timestamp(arrow::datatypes::TimeUnit::Nanosecond, None),
"Date32" => DataType::Date32, "Date32" => DataType::Date32,
_ => return Err(anyhow!("Unsupported type in config: {}", col_cfg.arrow_type)), _ => return Err(anyhow!("Unsupported type in config: {}", col_cfg.arrow_type)),
}; };
Ok(Field::new(&col_cfg.column_name, arrow_type, nullable)) Ok(Field::new(&col_cfg.column_name, arrow_type, nullable))
}).collect::<Result<Vec<Field>>>()?; }).collect::<Result<Vec<Field>>>()?;
Ok(Schema::new(schema_fields)) Ok(Schema::new(schema_fields))
} }
// -------------------------------------------------------------------------------- // --------------------------------------------------------------------------------
// --- 2. STATIC DISPATCH IMPLEMENTATION (The Speed Gain) --- // --- 2. STATIC DISPATCH IMPLEMENTATION (The Fast Parser) ---
// -------------------------------------------------------------------------------- // --------------------------------------------------------------------------------
// New Trait for all builders to implement append_null // Struct to hold the builders in a statically-known, fixed order (eliminates DynamicBuilder enum)
trait ColumnBuilderTrait {
fn append_null_to_self(&mut self);
}
// Implement the trait for the Boxed builders (which were in the DynamicBuilder enum)
impl ColumnBuilderTrait for Box<Int64Builder> { fn append_null_to_self(&mut self) { self.append_null(); } }
impl ColumnBuilderTrait for Box<Int32Builder> { fn append_null_to_self(&mut self) { self.append_null(); } }
impl ColumnBuilderTrait for Box<Float64Builder> { fn append_null_to_self(&mut self) { self.append_null(); } }
impl ColumnBuilderTrait for Box<Float32Builder> { fn append_null_to_self(&mut self) { self.append_null(); } }
impl ColumnBuilderTrait for Box<StringBuilder> { fn append_null_to_self(&mut self) { self.append_null(); } }
impl ColumnBuilderTrait for Box<BooleanBuilder> { fn append_null_to_self(&mut self) { self.append_null(); } }
impl ColumnBuilderTrait for Box<TimestampNanosecondBuilder> { fn append_null_to_self(&mut self) { self.append_null(); } }
impl ColumnBuilderTrait for Box<Date32Builder> { fn append_null_to_self(&mut self) { self.append_null(); } }
// New struct to hold the builders in a statically-known, fixed order
// Note: We use the exact types from the benchmark schema to simplify the MVP
struct SchemaParser { struct SchemaParser {
// Column 0: id
id: Box<Int64Builder>, id: Box<Int64Builder>,
// Column 1: uuid
uuid: Box<StringBuilder>, uuid: Box<StringBuilder>,
// Column 2: username
username: Box<StringBuilder>, username: Box<StringBuilder>,
// Column 3: score
score: Box<Float32Builder>, score: Box<Float32Builder>,
// Column 4: is_active
is_active: Box<BooleanBuilder>, is_active: Box<BooleanBuilder>,
// Column 5: last_login
last_login: Box<TimestampNanosecondBuilder>, last_login: Box<TimestampNanosecondBuilder>,
// Column 6: notes
notes: Box<StringBuilder>, notes: Box<StringBuilder>,
// Column 7: course_id
course_id: Box<Int32Builder>, course_id: Box<Int32Builder>,
// Column 8: start_date
start_date: Box<Date32Builder>, start_date: Box<Date32Builder>,
// Column 9: rating
rating: Box<Float64Builder>, rating: Box<Float64Builder>,
// Note: This struct MUST match the order of the query result.
} }
// Helper to construct and parse data using the static SchemaParser // Helper to construct and parse data using the static SchemaParser
@@ -259,7 +246,7 @@ async fn parse_data_with_schema(
Ok((rows_processed, record_batch)) Ok((rows_processed, record_batch))
} }
// The core streaming parser logic - generic over the SchemaParser struct // The core streaming parser logic
async fn parse_binary_stream_static( async fn parse_binary_stream_static(
mut stream: Pin<Box<CopyOutStream>>, mut stream: Pin<Box<CopyOutStream>>,
parser: &mut SchemaParser, parser: &mut SchemaParser,
@@ -276,13 +263,14 @@ async fn parse_binary_stream_static(
if !is_header_parsed { if !is_header_parsed {
if buffer.len() < 19 { continue 'stream_loop; } if buffer.len() < 19 { continue 'stream_loop; }
let mut header_cursor = Cursor::new(&buffer[..]); let mut header_cursor = Cursor::new(&buffer[..]);
parse_stream_header(&mut header_cursor)?; parse_stream_header(&mut header_cursor).context("Failed to parse stream header")?;
buffer.advance(19); buffer.advance(19);
is_header_parsed = true; is_header_parsed = true;
} }
'parsing_loop: loop { 'parsing_loop: loop {
let mut cursor = Cursor::new(&buffer[..]); let mut cursor = Cursor::new(&buffer[..]);
let safe_position = cursor.position();
let col_count = match cursor.read_i16::<BigEndian>() { let col_count = match cursor.read_i16::<BigEndian>() {
Ok(count) => count, Ok(count) => count,
@@ -302,6 +290,14 @@ async fn parse_binary_stream_static(
buffer.advance(bytes_consumed as usize); buffer.advance(bytes_consumed as usize);
} }
Err(e) if e.kind() == std::io::ErrorKind::UnexpectedEof => { Err(e) if e.kind() == std::io::ErrorKind::UnexpectedEof => {
cursor.set_position(safe_position);
// Copy remaining bytes back to the buffer for the next chunk
let remaining_slice = &buffer.as_ref()[safe_position as usize..];
let mut leftover_buffer_vec = Vec::new();
leftover_buffer_vec.extend_from_slice(remaining_slice);
buffer.clear();
buffer.extend_from_slice(&leftover_buffer_vec);
break 'parsing_loop; break 'parsing_loop;
} }
Err(e) => { Err(e) => {
@@ -320,7 +316,6 @@ async fn parse_binary_stream_static(
fn parse_stream_header(cursor: &mut Cursor<&[u8]>) -> Result<()> { fn parse_stream_header(cursor: &mut Cursor<&[u8]>) -> Result<()> {
// (Unchanged)
let mut magic_signature = [0u8; 11]; let mut magic_signature = [0u8; 11];
cursor.read_exact(&mut magic_signature).context("Failed to read magic signature")?; cursor.read_exact(&mut magic_signature).context("Failed to read magic signature")?;
if &magic_signature != b"PGCOPY\n\xff\r\n\0" { if &magic_signature != b"PGCOPY\n\xff\r\n\0" {
@@ -363,8 +358,8 @@ fn parse_row_static(
let len = cursor.read_i32::<BigEndian>()?; let len = cursor.read_i32::<BigEndian>()?;
if len == -1 { p.last_login.append_null() } else { if len == -1 { p.last_login.append_null() } else {
let pg_micros = cursor.read_i64::<BigEndian>()?; let pg_micros = cursor.read_i64::<BigEndian>()?;
// 10957 days between 1970 and 2000 => 946684800000000 micros // Optimization: Constant offset applied
let unix_micros = pg_micros + 946684800000000; let unix_micros = pg_micros + POSTGRES_EPOCH_MICROS_OFFSET;
p.last_login.append_value(unix_micros * 1000); p.last_login.append_value(unix_micros * 1000);
} }
@@ -380,7 +375,7 @@ fn parse_row_static(
let len = cursor.read_i32::<BigEndian>()?; let len = cursor.read_i32::<BigEndian>()?;
if len == -1 { p.start_date.append_null() } else { if len == -1 { p.start_date.append_null() } else {
let pg_days = cursor.read_i32::<BigEndian>()?; let pg_days = cursor.read_i32::<BigEndian>()?;
// 10957 days between 1970 and 2000 // Optimization: 10957 days between 1970 and 2000
p.start_date.append_value(pg_days + 10957); p.start_date.append_value(pg_days + 10957);
} }