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
9 changes: 4 additions & 5 deletions posthog/tasks/exports/image_exporter.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
import json
import os
import uuid
import tempfile
from datetime import timedelta
from typing import Literal, Optional

Expand Down Expand Up @@ -36,8 +37,6 @@

logger = structlog.get_logger(__name__)

TMP_DIR = "/tmp" # NOTE: Externalise this to ENV var

ScreenWidth = Literal[800, 1920]
CSSSelector = Literal[".InsightCard", ".ExportedInsight"]

Expand Down Expand Up @@ -85,8 +84,8 @@ def _export_to_png(exported_asset: ExportedAsset) -> None:
)

image_id = str(uuid.uuid4())
image_path = os.path.join(TMP_DIR, f"{image_id}.png")

with tempfile.TemporaryFile() as tmp:
pass
if not os.path.exists(TMP_DIR):
os.makedirs(TMP_DIR)

Expand All @@ -104,7 +103,7 @@ def _export_to_png(exported_asset: ExportedAsset) -> None:
wait_for_css_selector = ".InsightCard"
screenshot_width = 1920
else:
raise Exception(f"Export is missing required dashboard or insight ID")
raise Exception("Export is missing required dashboard or insight ID")

logger.info("exporting_asset", asset_id=exported_asset.id, render_url=url_to_render)

Expand Down
166 changes: 80 additions & 86 deletions posthog/warehouse/models/external_data_schema.py
Original file line number Diff line number Diff line change
Expand Up @@ -86,19 +86,19 @@ def update_incremental_field_last_value(self, last_value: Any) -> None:
if last_value_py is None:
return

if (
incremental_field_type == IncrementalFieldType.Integer
or incremental_field_type == IncrementalFieldType.Numeric
if incremental_field_type in (
IncrementalFieldType.Integer,
IncrementalFieldType.Numeric,
):
if isinstance(last_value_py, int | float):
last_value_json = last_value_py
elif isinstance(last_value_py, datetime):
last_value_json = last_value_py.isoformat()
else:
last_value_json = int(last_value_py)
elif (
incremental_field_type == IncrementalFieldType.DateTime
or incremental_field_type == IncrementalFieldType.Timestamp
elif incremental_field_type in (
IncrementalFieldType.DateTime,
IncrementalFieldType.Timestamp,
):
if isinstance(last_value_py, datetime):
last_value_json = last_value_py.isoformat()
Expand All @@ -110,12 +110,6 @@ def update_incremental_field_last_value(self, last_value: Any) -> None:
self.sync_type_config["incremental_field_last_value"] = last_value_json
self.save()

def soft_delete(self):
self.deleted = True
self.deleted_at = datetime.now()
self.save()


@database_sync_to_async
def asave_external_data_schema(schema: ExternalDataSchema) -> None:
schema.save()
Expand Down Expand Up @@ -273,20 +267,19 @@ def get_snowflake_schemas(
schema="information_schema",
role=role,
**auth_connect_args,
) as connection:
with connection.cursor() as cursor:
if cursor is None:
raise Exception("Can't create cursor to Snowflake")
) as connection, connection.cursor() as cursor:
if cursor is None:
raise Exception("Can't create cursor to Snowflake")

cursor.execute(
"SELECT table_name, column_name, data_type FROM information_schema.columns WHERE table_schema = %(schema)s ORDER BY table_name ASC",
{"schema": schema},
)
result = cursor.fetchall()
cursor.execute(
"SELECT table_name, column_name, data_type FROM information_schema.columns WHERE table_schema = %(schema)s ORDER BY table_name ASC",
{"schema": schema},
)
result = cursor.fetchall()

schema_list = defaultdict(list)
for row in result:
schema_list[row[0]].append((row[1], row[2]))
schema_list = defaultdict(list)
for row in result:
schema_list[row[0]].append((row[1], row[2]))

if file_name is not None:
os.unlink(file_name)
Expand All @@ -302,7 +295,7 @@ def filter_postgres_incremental_fields(columns: list[tuple[str, str]]) -> list[t
results.append((column_name, IncrementalFieldType.Timestamp))
elif type == "date":
results.append((column_name, IncrementalFieldType.Date))
elif type == "integer" or type == "smallint" or type == "bigint":
elif type in ("integer", "smallint", "bigint"):
results.append((column_name, IncrementalFieldType.Integer))

return results
Expand All @@ -311,45 +304,47 @@ def filter_postgres_incremental_fields(columns: list[tuple[str, str]]) -> list[t
def get_postgres_row_count(
host: str, port: str, database: str, user: str, password: str, schema: str, ssh_tunnel: SSHTunnel
) -> dict[str, int]:
import tempfile
def get_row_count(postgres_host: str, postgres_port: int):
connection = psycopg2.connect(
host=postgres_host,
port=postgres_port,
dbname=database,
user=user,
password=password,
sslmode="prefer",
connect_timeout=5,
sslrootcert="/tmp/no.txt",
sslcert="/tmp/no.txt",
sslkey="/tmp/no.txt",
)

try:
with connection.cursor() as cursor:
cursor.execute(
"SELECT tablename as table_name FROM pg_tables WHERE schemaname = %(schema)s",
{"schema": schema},
)
tables = cursor.fetchall()

if not tables:
return {}
with tempfile.TemporaryFile() as tmp:
connection = psycopg2.connect(
host=postgres_host,
port=postgres_port,
dbname=database,
user=user,
password=password,
sslmode="prefer",
connect_timeout=5,
sslrootcert=tmp.name,
sslcert=tmp.name,
sslkey=tmp.name,
)

counts = [
sql.SQL("SELECT {table_name} AS table_name, COUNT(*) AS row_count FROM {schema}.{table}").format(
table_name=sql.Literal(table[0]), schema=sql.Identifier(schema), table=sql.Identifier(table[0])
try:
with connection.cursor() as cursor:
cursor.execute(
"SELECT tablename as table_name FROM pg_tables WHERE schemaname = %(schema)s",
{"schema": schema},
)
for table in tables
]

union_counts = sql.SQL(" UNION ALL ").join(counts)
cursor.execute(union_counts)
row_count_result = cursor.fetchall()
row_counts = {row[0]: row[1] for row in row_count_result}
return row_counts
finally:
connection.close()
tables = cursor.fetchall()

if not tables:
return {}

counts = [
sql.SQL("SELECT {table_name} AS table_name, COUNT(*) AS row_count FROM {schema}.{table}").format(
table_name=sql.Literal(table[0]), schema=sql.Identifier(schema), table=sql.Identifier(table[0])
)
for table in tables
]

union_counts = sql.SQL(" UNION ALL ").join(counts)
cursor.execute(union_counts)
row_count_result = cursor.fetchall()
row_counts = {row[0]: row[1] for row in row_count_result}
return row_counts
finally:
connection.close()

if ssh_tunnel.enabled:
with ssh_tunnel.get_tunnel(host, int(port)) as tunnel:
Expand All @@ -364,32 +359,31 @@ def get_row_count(postgres_host: str, postgres_port: int):
def get_postgres_schemas(
host: str, port: str, database: str, user: str, password: str, schema: str, ssh_tunnel: SSHTunnel
) -> dict[str, list[tuple[str, str]]]:
import tempfile
def get_schemas(postgres_host: str, postgres_port: int):
connection = psycopg2.connect(
host=postgres_host,
port=postgres_port,
dbname=database,
user=user,
password=password,
sslmode="prefer",
connect_timeout=5,
sslrootcert="/tmp/no.txt",
sslcert="/tmp/no.txt",
sslkey="/tmp/no.txt",
)

with connection.cursor() as cursor:
cursor.execute(
"SELECT table_name, column_name, data_type FROM information_schema.columns WHERE table_schema = %(schema)s ORDER BY table_name ASC",
{"schema": schema},
with tempfile.TemporaryFile() as tmp:
connection = psycopg2.connect(
host=postgres_host,
port=postgres_port,
dbname=database,
user=user,
password=password,
sslmode="prefer",
connect_timeout=5,
)
result = cursor.fetchall()

schema_list = defaultdict(list)
for row in result:
schema_list[row[0]].append((row[1], row[2]))
with connection.cursor() as cursor:
cursor.execute(
"SELECT table_name, column_name, data_type FROM information_schema.columns WHERE table_schema = %(schema)s ORDER BY table_name ASC",
{"schema": schema},
)
result = cursor.fetchall()

connection.close()
schema_list = defaultdict(list)
for row in result:
schema_list[row[0]].append((row[1], row[2]))

connection.close()

return schema_list

Expand All @@ -413,7 +407,7 @@ def filter_mysql_incremental_fields(columns: list[tuple[str, str]]) -> list[tupl
results.append((column_name, IncrementalFieldType.Date))
elif type == "datetime":
results.append((column_name, IncrementalFieldType.DateTime))
elif type == "tinyint" or type == "smallint" or type == "mediumint" or type == "int" or type == "bigint":
elif type in ("tinyint", "smallint", "mediumint", "int", "bigint"):
results.append((column_name, IncrementalFieldType.Integer))

return results
Expand Down Expand Up @@ -476,9 +470,9 @@ def filter_mssql_incremental_fields(columns: list[tuple[str, str]]) -> list[tupl
type = type.lower()
if type == "date":
results.append((column_name, IncrementalFieldType.Date))
elif type == "datetime" or type == "datetime2" or type == "smalldatetime":
elif type in ("datetime", "datetime2", "smalldatetime"):
results.append((column_name, IncrementalFieldType.DateTime))
elif type == "tinyint" or type == "smallint" or type == "int" or type == "bigint":
elif type in ("tinyint", "smallint", "int", "bigint"):
results.append((column_name, IncrementalFieldType.Integer))

return results
Expand Down