Skip to content
Merged
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
2 changes: 2 additions & 0 deletions backend/protocol_rpc/app_lifespan.py
Original file line number Diff line number Diff line change
Expand Up @@ -26,6 +26,7 @@
)
from backend.protocol_rpc.transactions_parser import TransactionParser
from backend.protocol_rpc.configuration import GlobalConfiguration
from backend.protocol_rpc.explorer.query_runner import ExplorerQueryRunner
from backend.protocol_rpc.fastapi_rpc_router import FastAPIRPCRouter
from backend.protocol_rpc.message_handler.fastapi_handler import (
MessageHandler,
Expand Down Expand Up @@ -234,6 +235,7 @@ async def rpc_app_lifespan(app, settings: RPCAppSettings) -> AsyncIterator[RPCAp
)
db_manager = DatabaseSessionManager(settings.database_url)
set_database_manager(db_manager)
app.state.explorer_query_runner = ExplorerQueryRunner(db_manager)

logger.info("[STARTUP] Verifying database readiness and migrations")
_verify_database_ready(db_manager)
Expand Down
17 changes: 12 additions & 5 deletions backend/protocol_rpc/endpoints.py
Original file line number Diff line number Diff line change
Expand Up @@ -70,6 +70,7 @@
from backend.database_handler.snapshot_manager import SnapshotManager
from backend.node.base import Manager as GenVMManager
import asyncio
from starlette.concurrency import run_in_threadpool

# Limit concurrent GenVM executions on the jsonrpc path to prevent uvloop fd
# conflicts and DB pool exhaustion while calls hold request-scoped sessions.
Expand Down Expand Up @@ -1076,7 +1077,9 @@ async def get_contract_schema(
contract_address: str,
) -> dict:
try:
contract_snapshot = ContractSnapshot(contract_address, session)
contract_snapshot = await run_in_threadpool(
ContractSnapshot, contract_address, session
)
except ContractNotFoundError:
raise NotFoundError(
message=f"Contract {contract_address} not found",
Expand Down Expand Up @@ -1424,7 +1427,9 @@ async def _gen_call_with_validator(

# Create validator node
try:
contract_snapshot = ContractSnapshot(to_address, session)
contract_snapshot = await run_in_threadpool(
ContractSnapshot, to_address, session
)
except ContractNotFoundError:
raise NotFoundError(
message=f"Contract {to_address} not found",
Expand Down Expand Up @@ -1630,8 +1635,8 @@ async def eth_call(

# Check if this is a ConsensusData contract call that we should handle locally
# This should happen before early return to allow interception even without 'from'
consensus_data_result = handle_consensus_data_call(
transactions_processor, to_address, data
consensus_data_result = await run_in_threadpool(
handle_consensus_data_call, transactions_processor, to_address, data
)
if consensus_data_result is not None:
return consensus_data_result
Expand All @@ -1657,7 +1662,9 @@ async def eth_call(
)
as_validator = snapshot.nodes[0].validator
try:
target_contract_snapshot = ContractSnapshot(to_address, session)
target_contract_snapshot = await run_in_threadpool(
ContractSnapshot, to_address, session
)
except ContractNotFoundError:
raise NotFoundError(
message=f"Contract {to_address} not found",
Expand Down
69 changes: 69 additions & 0 deletions backend/protocol_rpc/explorer/query_runner.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,69 @@
"""Bound Explorer reads and close their sessions before returning to the API loop."""

import threading
import time
from typing import Any, Callable

from fastapi import HTTPException

from backend.database_handler.session_factory import DatabaseSessionManager

from . import queries


class ExplorerQueryRunner:
def __init__(
self,
db_manager: DatabaseSessionManager,
*,
max_concurrent: int = 2,
counts_ttl: float = 5.0,
) -> None:
self._db_manager = db_manager
self._slots = threading.BoundedSemaphore(max_concurrent)
self._counts_ttl = counts_ttl
self._counts_lock = threading.Lock()
self._cached_counts: tuple[float, dict] | None = None

@staticmethod
def _busy() -> HTTPException:
return HTTPException(
status_code=503,
detail="Explorer busy; retry shortly",
headers={"Retry-After": "1"},
)

def run(self, query: Callable[..., Any], *args: Any, **kwargs: Any) -> Any:
if not self._slots.acquire(blocking=False):
raise self._busy()
try:
# Query functions materialize their results. Closing here rolls back
# the read transaction in the same worker, without an event-loop hop.
with self._db_manager.open_session() as session:
return query(session, *args, **kwargs)
finally:
self._slots.release()

def counts(self) -> dict:
cached = self._cached_counts
if cached is not None and cached[0] > time.monotonic():
return dict(cached[1])

# Coalesce refreshes across Explorer server renders. A concurrent caller
# can use the previous counts while the single refresh is in progress.
if not self._counts_lock.acquire(blocking=False):
if cached is not None:
return dict(cached[1])
raise self._busy()
try:
cached = self._cached_counts
if cached is not None and cached[0] > time.monotonic():
return dict(cached[1])
counts = self.run(queries.get_stats_counts)
self._cached_counts = (
time.monotonic() + self._counts_ttl,
dict(counts),
)
return counts
finally:
self._counts_lock.release()
57 changes: 36 additions & 21 deletions backend/protocol_rpc/explorer/router.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,29 +2,37 @@

from typing import Annotated, Literal, Optional

from fastapi import APIRouter, Depends, HTTPException, Query
from sqlalchemy.orm import Session

from backend.protocol_rpc.dependencies import get_db_session
from fastapi import APIRouter, Depends, HTTPException, Query, Request

from . import queries
from .query_runner import ExplorerQueryRunner

explorer_router = APIRouter(prefix="/api/explorer", tags=["explorer"])


def get_query_runner(request: Request) -> ExplorerQueryRunner:
runner = getattr(request.app.state, "explorer_query_runner", None)
if runner is None:
raise HTTPException(status_code=503, detail="Explorer not initialized")
return runner


QueryRunner = Annotated[ExplorerQueryRunner, Depends(get_query_runner)]


# ---------------------------------------------------------------------------
# Stats
# ---------------------------------------------------------------------------


@explorer_router.get("/stats")
def get_stats(session: Annotated[Session, Depends(get_db_session)]):
return queries.get_stats(session)
def get_stats(runner: QueryRunner):
return runner.run(queries.get_stats)


@explorer_router.get("/stats/counts")
def get_stats_counts(session: Annotated[Session, Depends(get_db_session)]):
return queries.get_stats_counts(session)
def get_stats_counts(runner: QueryRunner):
return runner.counts()


# ---------------------------------------------------------------------------
Expand All @@ -34,7 +42,7 @@ def get_stats_counts(session: Annotated[Session, Depends(get_db_session)]):

@explorer_router.get("/transactions")
def get_transactions(
session: Annotated[Session, Depends(get_db_session)],
runner: QueryRunner,
page: int = Query(1, ge=1),
limit: int = Query(20, ge=1, le=100),
status: Optional[str] = None,
Expand All @@ -43,17 +51,24 @@ def get_transactions(
to_date: Optional[str] = None,
address: Optional[str] = None,
):
return queries.get_all_transactions_paginated(
session, page, limit, status, search, from_date, to_date, address
return runner.run(
queries.get_all_transactions_paginated,
page,
limit,
status,
search,
from_date,
to_date,
address,
)


@explorer_router.get("/transactions/{tx_hash}")
def get_transaction(
tx_hash: str,
session: Annotated[Session, Depends(get_db_session)],
runner: QueryRunner,
):
result = queries.get_transaction_with_relations(session, tx_hash)
result = runner.run(queries.get_transaction_with_relations, tx_hash)
if result is None:
raise HTTPException(status_code=404, detail="Transaction not found")
return result
Expand All @@ -66,11 +81,11 @@ def get_transaction(

@explorer_router.get("/validators")
def get_validators(
session: Annotated[Session, Depends(get_db_session)],
runner: QueryRunner,
search: Optional[str] = None,
limit: Optional[int] = Query(None, ge=1, le=100),
):
return queries.get_all_validators(session, search=search, limit=limit)
return runner.run(queries.get_all_validators, search=search, limit=limit)


# ---------------------------------------------------------------------------
Expand All @@ -81,9 +96,9 @@ def get_validators(
@explorer_router.get("/address/{address}")
def get_address(
address: str,
session: Annotated[Session, Depends(get_db_session)],
runner: QueryRunner,
):
result = queries.get_address_info(session, address)
result = runner.run(queries.get_address_info, address)
if result is None:
raise HTTPException(status_code=404, detail="Address not found")
return result
Expand All @@ -96,14 +111,14 @@ def get_address(

@explorer_router.get("/contracts")
def get_contracts(
session: Annotated[Session, Depends(get_db_session)],
runner: QueryRunner,
search: Optional[str] = None,
page: int = Query(1, ge=1),
limit: int = Query(20, ge=1, le=100),
sort_by: Optional[Literal["tx_count", "created_at", "updated_at"]] = None,
sort_order: Literal["asc", "desc"] = "desc",
):
return queries.get_all_states(session, search, page, limit, sort_by, sort_order)
return runner.run(queries.get_all_states, search, page, limit, sort_by, sort_order)


# ---------------------------------------------------------------------------
Expand All @@ -112,5 +127,5 @@ def get_contracts(


@explorer_router.get("/providers")
def get_providers(session: Annotated[Session, Depends(get_db_session)]):
return queries.get_all_providers(session)
def get_providers(runner: QueryRunner):
return runner.run(queries.get_all_providers)
26 changes: 22 additions & 4 deletions backend/protocol_rpc/message_handler/fastapi_handler.py
Original file line number Diff line number Diff line change
Expand Up @@ -30,10 +30,15 @@ def __init__(self, broadcast: Broadcast, config: GlobalConfiguration):
self.broadcast = broadcast
self.config = config
self.client_session_id = None
try:
self._loop = asyncio.get_running_loop()
except RuntimeError:
self._loop = None

def with_client_session(self, client_session_id: str):
new_msg_handler = MessageHandler(self.broadcast, self.config)
new_msg_handler.client_session_id = client_session_id
new_msg_handler._loop = self._loop
return new_msg_handler

def log_endpoint_info(self, func):
Expand Down Expand Up @@ -76,14 +81,27 @@ def _publish(self, channel: str, payload: dict[str, Any]) -> None:

message = json.dumps(payload)
try:
loop = asyncio.get_running_loop()
running_loop = asyncio.get_running_loop()
except RuntimeError:
return
running_loop = None

if not loop.is_running():
loop = self._loop or running_loop
if loop is None or not loop.is_running():
return
self._loop = loop

loop.create_task(self.broadcast.publish(channel=channel, message=message))
def publish():
loop.create_task(self.broadcast.publish(channel=channel, message=message))

if running_loop is loop:
publish()
else:
# Sync RPC handlers run in workers; Broadcast belongs to the app loop.
try:
loop.call_soon_threadsafe(publish)
except RuntimeError:
# The application may have shut down while a worker finished.
pass

def _socket_emit(self, log_event: LogEvent) -> None:
"""Emit a log event via broadcast channels.
Expand Down
13 changes: 11 additions & 2 deletions backend/protocol_rpc/rpc_endpoint_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,7 @@
from fastapi.dependencies.utils import get_dependant, solve_dependencies
from fastapi.requests import Request
from pydantic import BaseModel, ConfigDict
from starlette.concurrency import run_in_threadpool

from backend.protocol_rpc.exceptions import (
InternalError,
Expand Down Expand Up @@ -305,7 +306,7 @@ async def _call_endpoint(
call_kwargs = bound_arguments
if "msg_handler" in call_kwargs:
call_kwargs["msg_handler"] = session_logger
result = registered.dependant.call(**call_kwargs)
result = await self._invoke_handler(registered.dependant.call, call_kwargs)
if inspect.isawaitable(result):
result = await result
return result
Expand Down Expand Up @@ -366,11 +367,19 @@ async def _call_endpoint(
if "msg_handler" in call_kwargs:
call_kwargs["msg_handler"] = session_logger

result = registered.dependant.call(**call_kwargs)
result = await self._invoke_handler(registered.dependant.call, call_kwargs)
if inspect.isawaitable(result):
result = await result
return result

@staticmethod
async def _invoke_handler(handler: Any, kwargs: Dict[str, Any]) -> Any:
# A synchronous pool checkout must not block the loop that schedules
# other requests' session cleanup (and therefore returns connections).
if inspect.iscoroutinefunction(handler):
return handler(**kwargs)
return await run_in_threadpool(handler, **kwargs)

def _bind_rpc_arguments(
self,
registered: RegisteredEndpoint,
Expand Down
Loading
Loading