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
11 changes: 4 additions & 7 deletions shepherd_server/base_routes.py
Original file line number Diff line number Diff line change
Expand Up @@ -28,7 +28,7 @@
save_logs,
save_message,
)
from shepherd_utils.logger import QueryLogger, setup_logging
from shepherd_utils.logger import attach_query_handler, setup_logging
from shepherd_utils.otel import setup_tracer

setup_logging()
Expand Down Expand Up @@ -88,10 +88,9 @@ async def run_query(
# Set up logger
log_level = query.get("log_level") or settings.log_level
level_number = logging._nameToLevel[log_level]
log_handler = QueryLogger().log_handler
logger = logging.getLogger(f"shepherd.{query_id}")
logger.setLevel(level_number)
logger.addHandler(log_handler)
attach_query_handler(logger)

logger.info(f"Sending {query_id} to {target}")

Expand Down Expand Up @@ -309,9 +308,8 @@ async def callback(
# have, so the rejection / parse-error paths below can persist their logs
# too. The requested log level lives in the body -- which those paths never
# parse -- so leave the logger at its inherited default until we have it.
log_handler = QueryLogger().log_handler
logger = logging.getLogger(f"shepherd.{callback_id}")
logger.addHandler(log_handler)
attach_query_handler(logger)
max_bytes = settings.callback_max_request_size_bytes
raw = await _read_body_within_limit(request, max_bytes)
if raw is None:
Expand Down Expand Up @@ -448,10 +446,9 @@ async def get_query_response(
):
"""Get a query response."""
level_number = logging._nameToLevel["INFO"]
log_handler = QueryLogger().log_handler
logger = logging.getLogger("shepherd.get_query")
logger.setLevel(level_number)
logger.addHandler(log_handler)
attach_query_handler(logger)
response = await get_message(query_id, logger)
if response is None:
return JSONResponse(content={"error": "Not found"}, status_code=404)
Expand Down
7 changes: 7 additions & 0 deletions shepherd_utils/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -257,6 +257,13 @@ class Settings(BaseSettings):
# under pathological bursts; leftover ready callbacks are swept by the next
# drain iteration. 0 disables the cap (fold everything ready in one pass).
merge_max_fold: int = 25
# Cap on how many log entries are lifted out of a single callback message.
# Subservices report their retrieval work in the TRAPI ``logs`` list they
# post back, and those entries are folded into the query's log list. A
# chatty (or broken) one can return tens of thousands of them, which would
# then be stored per query and echoed in the final response, so keep only
# the first N. 0 disables the cap.
merge_max_callback_logs: int = 1000

# Monitor (dashboard) worker
monitor_port: int = 5440
Expand Down
110 changes: 76 additions & 34 deletions shepherd_utils/db.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,7 @@
from psycopg_pool import AsyncConnectionPool

from .config import settings
from .logger import get_query_handler

PG_RETRIES = 5

Expand Down Expand Up @@ -521,39 +522,75 @@ def save_message_sync(message_id: str, message: dict[str, Any]) -> None:
)


async def _append_logs(response_id: str, entries: List[dict]) -> None:
"""Append log entries to a query's list and (re)set the key's TTL.

Both in one round trip, so a crash can't leave the key without an
expiration.
"""
pipe = logs_db_client.pipeline()
pipe.rpush(response_id, *(orjson.dumps(entry) for entry in entries))
pipe.expire(response_id, settings.redis_ttl)
await pipe.execute()


async def _convert_legacy_logs(response_id: str) -> None:
"""Rewrite a pre-list logs key as a list of entries.

Logs used to be stored as one JSON array under this key, rewritten whole on
every flush. Appending to a list instead is atomic, so concurrent flushes
no longer clobber each other -- but a query mid-flight across the version
change still has the old blob. Replace it with the equivalent list; both
reads and writes fall back here on the type error and then carry on.

``GETDEL`` so that concurrent flushes can't each read the blob and push its
contents back a second time: exactly one of them comes away with it.
"""
blob = await logs_db_client.getdel(response_id)
entries = orjson.loads(blob) if blob else []
if entries:
await _append_logs(response_id, entries)


async def save_logs(
response_id: str,
logger: logging.Logger,
):
"""
Save logs from a worker to the db.

The query's logs are a Redis list that every flush appends to. Appending is
atomic, so the several producers a single query has -- the callback handler,
each worker stage, one merge task per callback -- can flush concurrently
without the read-modify-write of a whole-list rewrite dropping or doubling
anyone's entries.

Draining the handler is what keeps each record to a single copy: the logger
behind it is a process-wide singleton shared by every task for this query,
so anything left in the queue would be written again by the next flush.

Args:
response_id (str): UID for a query response
"""
# Drain before the first await: another flush for this query may interleave
# here, and each must come away with a disjoint set of records.
handler = get_query_handler(logger)
new_logs = handler.drain() if handler is not None else []
if not new_logs:
return
try:
existing_logs = await logs_db_client.get(response_id)
if existing_logs is None:
existing_logs = []
else:
existing_logs = orjson.loads(existing_logs)
# get log handler from logger
handler = next(
(
h
for h in logger.handlers
if getattr(h, "name", None) == "query_log_handler"
),
None,
)
if handler is not None:
new_logs = list(handler.contents())
new_logs.reverse()
existing_logs.extend(new_logs)
await logs_db_client.set(
response_id, orjson.dumps(existing_logs), ex=settings.redis_ttl
)
try:
await _append_logs(response_id, new_logs)
except redis.ResponseError:
# Key still holds the single-JSON-blob format this used to write --
# a query that was already in flight when this version rolled out.
# Convert it in place so its earlier logs survive.
await _convert_legacy_logs(response_id)
await _append_logs(response_id, new_logs)
except Exception as e:
# Put them back so the next flush retries rather than dropping them --
# they're no longer anywhere else now that the queue has been drained.
handler.ingest(new_logs)
logger.error(f"Failed to save logs for response {response_id}: {e}")


Expand All @@ -568,23 +605,28 @@ async def get_logs(
response_id (str): UID for a query response
"""
try:
logs = await logs_db_client.get(response_id)
if logs is not None:
logs = orjson.loads(logs)
# The stored list reflects the order in which workers *flushed*
# their logs, not the order events happened: a callback can arrive
# and be merged (and flushed) before the lookup that dispatched the
# query finishes and flushes its own logs. Every entry carries an
# ISO8601 UTC timestamp, so sort by it to present the logs in
# chronological order. Stable sort keeps same-timestamp entries in
# their original relative order.
logs.sort(key=lambda entry: entry.get("timestamp", ""))
return logs
else:
try:
entries = await logs_db_client.lrange(response_id, 0, -1)
except redis.ResponseError:
# Pre-list format (see ``_convert_legacy_logs``).
await _convert_legacy_logs(response_id)
entries = await logs_db_client.lrange(response_id, 0, -1)
if not entries:
logger.error(f"Failed to get logs for response {response_id}")
return []
logs = [orjson.loads(entry) for entry in entries]
# The stored list reflects the order in which workers *flushed*
# their logs, not the order events happened: a callback can arrive
# and be merged (and flushed) before the lookup that dispatched the
# query finishes and flushes its own logs. Every entry carries an
# ISO8601 UTC timestamp, so sort by it to present the logs in
# chronological order. Stable sort keeps same-timestamp entries in
# their original relative order.
logs.sort(key=lambda entry: entry.get("timestamp", ""))
return logs
except Exception as e:
logger.error(f"Failed to get logs from query: {e}")
return []


async def add_callback_id(
Expand Down
42 changes: 42 additions & 0 deletions shepherd_utils/logger.py
Original file line number Diff line number Diff line change
Expand Up @@ -63,6 +63,22 @@ def ingest(self, entries):
for entry in entries:
self.log_queue.appendleft(entry)

def drain(self):
"""Pop everything queued and return it oldest-first.

Reading is destructive on purpose. These handlers hang off loggers that
``logging.getLogger`` hands back as process-wide singletons, so the same
queue is flushed repeatedly -- once per task a worker runs for the same
query. Leaving entries behind meant every flush re-persisted every
record the queue had accumulated since the process started, which is how
one log line ended up in a query's logs several times over. Callers that
fail to persist what they drained put it back with ``ingest``.
"""
entries = list(self.log_queue)
entries.reverse()
self.log_queue.clear()
return entries


# Create unique logger for each query
# https://stackoverflow.com/a/37967421
Expand All @@ -80,6 +96,32 @@ def log_handler(self):
return self._log_handler


def get_query_handler(logger: logging.Logger):
"""Return the query log handler attached to ``logger``, or None."""
return next(
(h for h in logger.handlers if getattr(h, "name", None) == "query_log_handler"),
None,
)


def attach_query_handler(logger: logging.Logger) -> QueryLogHandler:
"""Give ``logger`` a query log handler, reusing one it already has.

``logging.getLogger`` returns the same object for a given name for the life
of the process, so callers that build "a logger per task" or "per request"
keep landing on one shared logger. Adding a fresh handler each time stacked
them up: every record was then queued once per attached handler, the list
grew without bound, and ``save_logs`` -- which flushes whichever handler
comes first -- kept re-reading the oldest queue. One handler per logger,
drained on flush, stores each record exactly once.
"""
handler = get_query_handler(logger)
if handler is None:
handler = QueryLogger().log_handler
logger.addHandler(handler)
return handler


def get_logging_config():
"""
Returns logging configuration.
Expand Down
11 changes: 6 additions & 5 deletions shepherd_utils/shared.py
Original file line number Diff line number Diff line change
Expand Up @@ -25,7 +25,7 @@
from .config import settings
from .db import initialize_db, save_logs
from .heartbeat import Heartbeat
from .logger import QueryLogger, setup_logging
from .logger import attach_query_handler, setup_logging
from .reclaim import reclaim_orphaned

# Cap each per-stream duration queue so a stopped monitor can't OOM the broker.
Expand Down Expand Up @@ -325,13 +325,15 @@ def _build_task_context(
level_number: int,
) -> Tuple[Context, logging.Logger]:
"""Build the per-task logger and otel context for a fetched/reclaimed task."""
log_handler = QueryLogger().log_handler
task_logger = logging.getLogger(
f"shepherd.{stream}.{consumer}.{ara_task[1]['query_id']}"
)
task_log_level = int(ara_task[1].get("log_level", level_number))
task_logger.setLevel(task_log_level)
task_logger.addHandler(log_handler)
# One handler per logger: this logger is shared by every task this worker
# runs for the query, and stacking a handler per task made each flush
# re-persist the previous tasks' records.
attach_query_handler(task_logger)
task_logger.debug(f"Doing task {ara_task}")
ctx = extract(json.loads(ara_task[1].get("otel", "{}")))
# Stamp the task payload with our delivery time so wrap_up_task /
Expand Down Expand Up @@ -481,10 +483,9 @@ async def get_tasks(
"""
# Set up logger
level_number = logging._nameToLevel[settings.log_level]
log_handler = QueryLogger().log_handler
worker_logger = logging.getLogger(f"shepherd.{stream}.{consumer}")
worker_logger.setLevel(level_number)
worker_logger.addHandler(log_handler)
attach_query_handler(worker_logger)
# allow ops to tune concurrency per Deployment without a code change
task_limit = _resolve_task_limit(stream, task_limit, worker_logger)
# initialize opens the db connection
Expand Down
Loading
Loading