33import os
44import shlex
55import shutil
6- import subprocess # nosec
6+ import subprocess
77from typing import cast , Dict , List
88
9- from psycopg2 . _psycopg import cursor
9+ from psycopg import Cursor , sql
1010from sqlalchemy import text
1111from sqlalchemy .engine import create_engine , Engine
1212from sqlalchemy .engine .url import URL
13+ from typing_extensions import LiteralString
1314
1415from databudgie .adapter .base import Adapter , QueryResult
1516from databudgie .output import Console , default_console
@@ -46,24 +47,29 @@ def export_query(self, query: str) -> QueryResult:
4647 result = QueryResult ()
4748 with result .binary_buffer () as buffer :
4849 with contextlib .closing (engine .raw_connection ()) as conn :
49- with cast (cursor , conn .cursor ()) as cursor_ :
50- copy = f"COPY ({ query } ) TO STDOUT CSV HEADER"
50+ with cast (Cursor , conn .cursor ()) as cursor :
51+ statement = sql . SQL ( cast ( LiteralString , f"COPY ({ query } ) TO STDOUT CSV HEADER" ))
5152
52- cursor_ .copy_expert (copy , buffer )
53- result .row_count = cursor_ .rowcount
53+ with cursor .copy (statement ) as copy :
54+ while data := copy .read ():
55+ buffer .write (data )
56+ result .row_count = cursor .rowcount
5457
5558 return result
5659
5760 def import_csv (self , csv_file : io .TextIOBase , table : str ):
5861 engine : Engine = cast (Engine , self .session .get_bind ())
5962
6063 # Reading the header line from the buffer removes it for the ingest
61- columns : List [str ] = [f'"{ c } "' for c in csv_file .readline ().strip ().split ("," )]
62- copy = "COPY {table} ({columns}) FROM STDIN CSV" .format (table = table , columns = "," .join (columns ))
64+ columns = [sql .Identifier (c ) for c in csv_file .readline ().strip ().split ("," )]
65+ statement = sql .SQL ("COPY {table} ({columns}) FROM STDIN CSV" ).format (
66+ table = sql .Identifier (* table .split ("." )), columns = sql .SQL (", " ).join (columns )
67+ )
6368
6469 with contextlib .closing (engine .raw_connection ()) as conn :
65- with cast (cursor , conn .cursor ()) as cursor_ :
66- cursor_ .copy_expert (copy , csv_file )
70+ with cast (Cursor , conn .cursor ()) as cursor :
71+ with cursor .copy (statement ) as copy :
72+ copy .write (csv_file .read ())
6773 conn .commit ()
6874
6975 def export_schema_ddl (self , name : str , console : Console = default_console ) -> bytes :
@@ -73,7 +79,10 @@ def export_schema_ddl(self, name: str, console: Console = default_console) -> by
7379
7480 url = self .session .connection ().engine .url
7581 result = pg_dump (url , f"--schema-only --schema={ name } --exclude-table={ name } .*" )
76- return result .replace (f"CREATE SCHEMA { name } ;" .encode (), f"CREATE SCHEMA IF NOT EXISTS { name } ;" .encode ())
82+ return result .replace (
83+ f"CREATE SCHEMA { name } ;" .encode (),
84+ f"CREATE SCHEMA IF NOT EXISTS { name } ;" .encode (),
85+ )
7786
7887 def export_table_ddl (self , table_name : str , console : Console = default_console ):
7988 if not shutil .which ("pg_dump" ):
@@ -205,7 +214,8 @@ def collect_table_dependencies(self, table_op: TableOp, console: Console = defau
205214 )
206215
207216 results = self .session .execute (
208- collect_tables , params = {"schema" : table_op .schema , "table_name" : table_op .table_name }
217+ collect_tables ,
218+ params = {"schema" : table_op .schema , "table_name" : table_op .table_name },
209219 )
210220
211221 return [row [0 ] for row in results ]
@@ -235,10 +245,18 @@ def collect_table_sequences(self) -> Dict[str, List[str]]:
235245 return result
236246
237247 def collect_sequence_value (self , sequence_name : str ) -> int :
238- return cast (int , self .session .execute (text (f"SELECT last_value from { sequence_name } " )).scalar ()) # noqa: S608
248+ return cast (
249+ int ,
250+ self .session .execute (
251+ text (f"SELECT last_value from { sequence_name } " ) # noqa: S608
252+ ).scalar (),
253+ )
239254
240255 def restore_sequence_value (self , sequence_name : str , value : int ) -> int :
241- return cast (int , self .session .execute (text (f"SELECT setval('{ sequence_name } ', { value } )" )).scalar ())
256+ return cast (
257+ int ,
258+ self .session .execute (text (f"SELECT setval('{ sequence_name } ', { value } )" )).scalar (),
259+ )
242260
243261
244262def pg_dump (url : URL , rest : str = "" , no_comments = True , clean = True ) -> bytes :
@@ -254,8 +272,11 @@ def pg_dump(url: URL, rest: str = "", no_comments=True, clean=True) -> bytes:
254272 command = shlex .split (raw_command )
255273
256274 try :
257- result = subprocess .run ( # nosec
258- command , capture_output = True , env = {** os .environ , "PGPASSWORD" : str (url .password or "" )}, check = True
275+ result = subprocess .run ( # noqa: S603
276+ command ,
277+ capture_output = True ,
278+ env = {** os .environ , "PGPASSWORD" : str (url .password or "" )},
279+ check = True ,
259280 )
260281 except subprocess .CalledProcessError as e :
261282 raise RuntimeError (e .stderr )
0 commit comments