from __future__ import annotations
import threading
from time import perf_counter
from typing import TYPE_CHECKING
from bottle import Bottle
from bottle import SimpleTemplate
from optuna.storages import RDBStorage
from optuna_dashboard._app import BottleViewReturn
from sqlalchemy import event
if TYPE_CHECKING:
from sqlalchemy.engine.base import Engine
sql_queries_lock = threading.Lock()
sql_queries: dict[str, tuple[int, list[float]]] = {}
sql_queries_template = SimpleTemplate(
"""
SQL Profiler - Optuna Dashboard
SQL Profiler
Sort by Total Time
| Total Time (s) |
Query Count |
Statement |
%for query in sort_by_total:
| {{ query[2] }}
| {{ query[1] }}
| {{ query[0] }}
|
%end
Sort by Count
| Query Count |
Total Time (s) |
Statement |
%for query in sort_by_count:
| {{ query[1] }}
| {{ query[2] }}
| {{ query[0] }}
|
%end
""" # noqa: E501
)
class EngineDebuggingSignalEvents:
"""Sets up handlers for two events that let us track the execution time of
queries."""
def __init__(self, engine: "Engine") -> None:
self.engine = engine
self.query_start_time = perf_counter()
def register(self) -> None:
event.listen(self.engine, "before_cursor_execute", self.before_cursor_execute)
event.listen(self.engine, "after_cursor_execute", self.after_cursor_execute)
def before_cursor_execute( # type: ignore
self, conn, cursor, statement, parameters, context, executemany
) -> None:
self.query_start_time = perf_counter()
def after_cursor_execute( # type: ignore
self, conn, cursor, stmt, parameters, context, executemany
) -> None:
duration = perf_counter() - self.query_start_time
with sql_queries_lock:
registered = stmt in sql_queries
sql_queries[stmt] = (
sql_queries[stmt][0] + 1 if registered else 1,
sql_queries[stmt][1] + [duration] if registered else [duration],
)
def register_profiler_view(app: Bottle, storage: RDBStorage) -> Bottle:
EngineDebuggingSignalEvents(storage.engine).register()
@app.get("/sql-profiler")
def profile_sql_queries() -> BottleViewReturn:
global sql_queries
with sql_queries_lock:
summary = [
(stmt, count, f"{sum(durations):.4f}", sum(durations))
for stmt, (count, durations) in sql_queries.items()
]
sort_by_total = sorted(summary, key=lambda r: r[3], reverse=True)
sort_by_count = sorted(summary, key=lambda r: r[1], reverse=True)
res = sql_queries_template.render(
sort_by_total=sort_by_total[:5], sort_by_count=sort_by_count[:5]
)
sql_queries = {}
return res
return app