diff --git a/shepherd_server/base_routes.py b/shepherd_server/base_routes.py index 6abc5d7..579ad67 100644 --- a/shepherd_server/base_routes.py +++ b/shepherd_server/base_routes.py @@ -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() @@ -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}") @@ -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: @@ -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) diff --git a/shepherd_utils/config.py b/shepherd_utils/config.py index 78e9907..d8d85a6 100644 --- a/shepherd_utils/config.py +++ b/shepherd_utils/config.py @@ -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 diff --git a/shepherd_utils/db.py b/shepherd_utils/db.py index dfdff43..956e4df 100644 --- a/shepherd_utils/db.py +++ b/shepherd_utils/db.py @@ -14,6 +14,7 @@ from psycopg_pool import AsyncConnectionPool from .config import settings +from .logger import get_query_handler PG_RETRIES = 5 @@ -521,6 +522,36 @@ 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, @@ -528,32 +559,38 @@ async def save_logs( """ 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}") @@ -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( diff --git a/shepherd_utils/logger.py b/shepherd_utils/logger.py index e43957b..8fb62e8 100644 --- a/shepherd_utils/logger.py +++ b/shepherd_utils/logger.py @@ -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 @@ -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. diff --git a/shepherd_utils/shared.py b/shepherd_utils/shared.py index b97f9ff..3a34bd8 100644 --- a/shepherd_utils/shared.py +++ b/shepherd_utils/shared.py @@ -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. @@ -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 / @@ -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 diff --git a/tests/unit/test_db.py b/tests/unit/test_db.py index 436882a..d28ff39 100644 --- a/tests/unit/test_db.py +++ b/tests/unit/test_db.py @@ -6,6 +6,7 @@ ``postgres_mock`` from the conftest. """ +import asyncio import io import logging @@ -13,6 +14,7 @@ import pytest import zstandard +import shepherd_utils.db as db_module from shepherd_utils.config import settings from shepherd_utils.db import ( ResponseTooLargeError, @@ -30,7 +32,7 @@ save_message, save_message_sync, ) -from shepherd_utils.logger import QueryLogger +from shepherd_utils.logger import attach_query_handler logger = logging.getLogger(__name__) @@ -192,55 +194,147 @@ def test_get_message_sync_raises_keyerror_for_missing(mocker): get_message_sync("missing-sid") +def _query_logger(name): + """A logger with nothing but a fresh query log handler on it.""" + sub_logger = logging.getLogger(name) + sub_logger.handlers.clear() + sub_logger.setLevel(logging.DEBUG) + attach_query_handler(sub_logger) + return sub_logger + + @pytest.mark.asyncio async def test_save_logs_appends_query_log_handler_records(redis_mock): """save_logs reads logs from a QueryLogHandler attached to the logger and - persists them (newest-first reversed) into the logs db.""" - handler = QueryLogger().log_handler - sub_logger = logging.getLogger("test.save_logs.appends") - sub_logger.handlers.clear() - sub_logger.addHandler(handler) - sub_logger.setLevel(logging.DEBUG) + persists them (oldest-first) into the logs db.""" + sub_logger = _query_logger("test.save_logs.appends") sub_logger.info("first message") sub_logger.info("second message") try: await save_logs("resp-1", sub_logger) finally: - sub_logger.removeHandler(handler) + sub_logger.handlers.clear() - raw = await redis_mock["logs"].get("resp-1") - assert raw is not None - logs = orjson.loads(raw) - messages = [entry["message"] for entry in logs] - # Insertion order: handler emits to a deque (appendleft), reversed in - # save_logs, so logs end up oldest-first. - assert messages == ["first message", "second message"] + logs = await get_logs("resp-1", logger) + # Insertion order: handler emits to a deque (appendleft), reversed on drain, + # so logs end up oldest-first. + assert [entry["message"] for entry in logs] == ["first message", "second message"] @pytest.mark.asyncio async def test_save_logs_extends_existing_logs(redis_mock): - """A pre-existing log array in redis is preserved and extended.""" + """Each flush appends; whatever earlier flushes stored is preserved.""" + earlier = _query_logger("test.save_logs.earlier") + earlier.info("from-earlier") + await save_logs("resp-2", earlier) + earlier.handlers.clear() + + sub_logger = _query_logger("test.save_logs.extends") + sub_logger.info("new entry") + try: + await save_logs("resp-2", sub_logger) + finally: + sub_logger.handlers.clear() + + logs = await get_logs("resp-2", logger) + assert [entry["message"] for entry in logs] == ["from-earlier", "new entry"] + + +@pytest.mark.asyncio +async def test_save_logs_does_not_rewrite_already_flushed_records(redis_mock): + """The regression this guards: a worker runs several tasks for one query and + they all share a logger, so a flush that left its records in the handler had + them written again by the next flush -- once more per task.""" + sub_logger = _query_logger("test.save_logs.no_dupes") + try: + sub_logger.info("callback one retrieval log") + await save_logs("resp-dupes", sub_logger) + sub_logger.info("callback two retrieval log") + await save_logs("resp-dupes", sub_logger) + # A flush with nothing new to say stores nothing at all. + await save_logs("resp-dupes", sub_logger) + finally: + sub_logger.handlers.clear() + + logs = await get_logs("resp-dupes", logger) + assert [entry["message"] for entry in logs] == [ + "callback one retrieval log", + "callback two retrieval log", + ] + + +@pytest.mark.asyncio +async def test_save_logs_keeps_records_when_the_write_fails(redis_mock, mocker): + """Draining is destructive, so a failed write has to put the records back -- + they're nowhere else -- for the next flush to retry.""" + sub_logger = _query_logger("test.save_logs.retry") + try: + sub_logger.info("only copy") + mocker.patch.object( + db_module.logs_db_client, "pipeline", side_effect=RuntimeError("redis down") + ) + await save_logs("resp-retry", sub_logger) + assert await get_logs("resp-retry", logger) == [] + + mocker.stopall() + await save_logs("resp-retry", sub_logger) + finally: + sub_logger.handlers.clear() + + assert "only copy" in [ + entry["message"] for entry in await get_logs("resp-retry", logger) + ] + + +@pytest.mark.asyncio +async def test_save_logs_concurrent_flushes_all_land(redis_mock): + """Two producers flushing the same query at once is normal (the callback + handler and a merge task, say). Appends are atomic, so neither loses.""" + loggers = [] + for i in range(5): + sub_logger = _query_logger(f"test.save_logs.concurrent.{i}") + sub_logger.info(f"entry {i}") + loggers.append(sub_logger) + try: + await asyncio.gather(*(save_logs("resp-concurrent", lg) for lg in loggers)) + finally: + for sub_logger in loggers: + sub_logger.handlers.clear() + + logs = await get_logs("resp-concurrent", logger) + assert sorted(entry["message"] for entry in logs) == [ + f"entry {i}" for i in range(5) + ] + + +@pytest.mark.asyncio +async def test_save_logs_converts_pre_list_logs_key(redis_mock): + """A query already in flight when this version rolls out has its logs stored + in the old whole-blob format; it's converted rather than lost.""" existing = [ {"message": "from-earlier", "timestamp": "2024-01-01T00:00:00", "level": "INFO"} ] - await redis_mock["logs"].set("resp-2", orjson.dumps(existing)) + await redis_mock["logs"].set("resp-legacy", orjson.dumps(existing)) - handler = QueryLogger().log_handler - sub_logger = logging.getLogger("test.save_logs.extends") - sub_logger.handlers.clear() - sub_logger.addHandler(handler) - sub_logger.setLevel(logging.DEBUG) - sub_logger.info("new entry") + sub_logger = _query_logger("test.save_logs.legacy") try: - await save_logs("resp-2", sub_logger) + sub_logger.info("new entry") + await save_logs("resp-legacy", sub_logger) finally: - sub_logger.removeHandler(handler) + sub_logger.handlers.clear() - raw = await redis_mock["logs"].get("resp-2") - logs = orjson.loads(raw) + logs = await get_logs("resp-legacy", logger) assert [entry["message"] for entry in logs] == ["from-earlier", "new entry"] +@pytest.mark.asyncio +async def test_get_logs_reads_pre_list_logs_key(redis_mock): + """Same, for a query that finishes before anything flushes to it again.""" + stored = [{"message": "hello", "timestamp": "ts", "level": "INFO"}] + await redis_mock["logs"].set("resp-legacy-read", orjson.dumps(stored)) + assert await get_logs("resp-legacy-read", logger) == stored + + @pytest.mark.asyncio async def test_get_logs_returns_empty_list_when_missing(redis_mock): """Reading logs for an unknown response id should return an empty list.""" diff --git a/tests/unit/test_logger.py b/tests/unit/test_logger.py index df97bfd..a4a812a 100644 --- a/tests/unit/test_logger.py +++ b/tests/unit/test_logger.py @@ -6,7 +6,9 @@ from shepherd_utils.logger import ( QueryLogger, ReasonerLogEntryFormatter, + attach_query_handler, get_logging_config, + get_query_handler, get_worker_logger, ) @@ -67,6 +69,68 @@ def test_query_logger_handler_named_query_log_handler(): assert handler.name == "query_log_handler" +def test_drain_empties_the_queue_and_returns_oldest_first(): + """Reading is destructive: the same records must not be handed out twice.""" + handler = QueryLogger().log_handler + sub_logger = logging.getLogger("test_query_logger.drain") + sub_logger.handlers.clear() + sub_logger.addHandler(handler) + sub_logger.setLevel(logging.DEBUG) + try: + sub_logger.info("first") + sub_logger.info("second") + assert [entry["message"] for entry in handler.drain()] == ["first", "second"] + assert handler.drain() == [] + sub_logger.info("third") + assert [entry["message"] for entry in handler.drain()] == ["third"] + finally: + sub_logger.removeHandler(handler) + + +def test_drained_records_can_be_put_back(): + """A failed flush returns its records to the queue for the next attempt.""" + handler = QueryLogger().log_handler + sub_logger = logging.getLogger("test_query_logger.putback") + sub_logger.handlers.clear() + sub_logger.addHandler(handler) + sub_logger.setLevel(logging.DEBUG) + try: + sub_logger.info("first") + sub_logger.info("second") + entries = handler.drain() + handler.ingest(entries) + assert [entry["message"] for entry in handler.drain()] == ["first", "second"] + finally: + sub_logger.removeHandler(handler) + + +def test_attach_query_handler_does_not_stack_handlers(): + """``logging.getLogger`` hands back one object per name, so a caller that + "makes a logger per task" keeps getting the same one. Attaching must be + idempotent -- stacked handlers queued every record once per handler and made + each flush re-persist the earlier tasks' records.""" + sub_logger = logging.getLogger("test_query_logger.attach") + sub_logger.handlers.clear() + sub_logger.setLevel(logging.DEBUG) + try: + first = attach_query_handler(sub_logger) + second = attach_query_handler(sub_logger) + assert first is second + assert sub_logger.handlers == [first] + assert get_query_handler(sub_logger) is first + + sub_logger.info("once") + assert [entry["message"] for entry in first.drain()] == ["once"] + finally: + sub_logger.handlers.clear() + + +def test_get_query_handler_returns_none_without_one(): + sub_logger = logging.getLogger("test_query_logger.none") + sub_logger.handlers.clear() + assert get_query_handler(sub_logger) is None + + def test_query_logger_respects_maxlen(): """When a maxlen is set, oldest records get dropped.""" ql = QueryLogger(maxlen=2) diff --git a/tests/unit/test_merge_message_batch.py b/tests/unit/test_merge_message_batch.py index 2cc7bd1..5c7d0e7 100644 --- a/tests/unit/test_merge_message_batch.py +++ b/tests/unit/test_merge_message_batch.py @@ -3,6 +3,8 @@ Covers: - ``merge_messages_by_ids`` folds multiple callbacks with a single load/save and is equivalent to merging them one at a time (the pre-existing behavior). +- The KG retrieval logs a callback carries back are handed to the parent + instead of being dropped with the rest of the callback message. - ``merge_messages_by_id`` still works as a one-callback delegate. - The per-query "ready callback" index helpers in ``shepherd_utils.db``. """ @@ -24,10 +26,12 @@ response_1, response_2, ) +from shepherd_utils.config import settings from workers.merge_message.worker import ( merge_messages, merge_messages_by_id, merge_messages_by_ids, + take_callback_logs, ) logger = logging.getLogger(__name__) @@ -144,6 +148,166 @@ def test_merge_messages_by_ids_child_handler_not_leaked(mocker): ) +def _callback_with_logs(logs): + """A callback message carrying the log entries a subservice returned.""" + callback = copy.deepcopy(response_2) + callback["logs"] = logs + return callback + + +def test_take_callback_logs_returns_and_clears_entries(): + """The subservice's entries come back out tagged with the callback id they + arrived under, and the field is blanked on the message so the log store + stays the single source of the final logs.""" + entries = [ + { + "timestamp": "2024-01-01T00:00:00+00:00", + "level": "INFO", + "message": "Calling KP infores:example", + }, + { + "timestamp": "2024-01-01T00:00:01+00:00", + "level": "WARNING", + "message": "KP infores:example timed out", + }, + ] + callback = _callback_with_logs(copy.deepcopy(entries)) + + taken = take_callback_logs(callback, "c1", logging.INFO, logger) + + assert taken == [ + {**entry, "message": f"[c1] {entry['message']}"} for entry in entries + ] + assert callback["logs"] == [] + + +def test_take_callback_logs_tags_entries_without_a_message(): + """A LogEntry can carry only a code; it still gets the callback id.""" + callback = _callback_with_logs([{"level": "INFO", "code": "KPTimeout"}]) + + taken = take_callback_logs(callback, "c1", logging.INFO, logger) + + assert taken == [{"level": "INFO", "code": "KPTimeout", "message": "[c1]"}] + + +def test_take_callback_logs_tags_each_callback_separately(): + """One query fans out into many retrievals whose logs all land in the same + list -- the tag is what says which retrieval a line came from.""" + first = take_callback_logs( + _callback_with_logs([{"level": "INFO", "message": "found 3 edges"}]), + "cb-one", + logging.INFO, + logger, + ) + second = take_callback_logs( + _callback_with_logs([{"level": "INFO", "message": "found 3 edges"}]), + "cb-two", + logging.INFO, + logger, + ) + + assert first[0]["message"] == "[cb-one] found 3 edges" + assert second[0]["message"] == "[cb-two] found 3 edges" + + +def test_take_callback_logs_filters_below_requested_level(): + """A subservice that reports at DEBUG shouldn't flood an INFO query.""" + callback = _callback_with_logs( + [ + {"level": "DEBUG", "message": "chatty"}, + {"level": "INFO", "message": "useful"}, + {"message": "no level at all"}, + ] + ) + + taken = take_callback_logs(callback, "c1", logging.INFO, logger) + + assert [entry["message"] for entry in taken] == [ + "[c1] useful", + "[c1] no level at all", + ] + + +def test_take_callback_logs_handles_malformed_logs(): + """Missing, null, non-list, and non-dict logs are all survivable.""" + assert take_callback_logs({}, "c1", logging.INFO, logger) == [] + assert take_callback_logs({"logs": None}, "c1", logging.INFO, logger) == [] + assert take_callback_logs({"logs": "nope"}, "c1", logging.INFO, logger) == [] + assert take_callback_logs( + {"logs": ["bare string"]}, "c1", logging.INFO, logger + ) == [{"message": "[c1] bare string", "level": "INFO"}] + + +def test_take_callback_logs_caps_entry_count(mocker): + """A subservice dumping tens of thousands of entries is truncated rather + than stored (and echoed back) in full.""" + mocker.patch.object(settings, "merge_max_callback_logs", 2) + callback = _callback_with_logs( + [{"level": "INFO", "message": f"log {i}"} for i in range(5)] + ) + + taken = take_callback_logs(callback, "c1", logging.INFO, logger) + + assert [entry["message"] for entry in taken] == ["[c1] log 0", "[c1] log 1"] + + +def test_merge_messages_by_ids_returns_callback_logs(mocker): + """The logs the KG retrieval sent back with each callback are returned to + the parent (which folds them into the query's log list) rather than being + dropped when the callback is merged into the response.""" + query_graph = response_1["message"]["query_graph"] + _patch_sync_store(mocker) + from shepherd_utils.db import save_message_sync + + save_message_sync("qid", {"message": {"query_graph": copy.deepcopy(query_graph)}}) + save_message_sync("rid", generate_response()) + save_message_sync( + "c1", _callback_with_logs([{"level": "INFO", "message": "c1 retrieval log"}]) + ) + save_message_sync( + "c2", _callback_with_logs([{"level": "ERROR", "message": "c2 retrieval log"}]) + ) + + merged, log_entries = merge_messages_by_ids( + "test_ara", "qid", "rid", ["c1", "c2"], logging.INFO + ) + + assert merged == ["c1", "c2"] + messages = [entry.get("message") for entry in log_entries] + # Tagged with the callback each came back on, so a line can be traced to + # the retrieval that emitted it. + assert "[c1] c1 retrieval log" in messages + assert "[c2] c2 retrieval log" in messages + # Retrieval logs are oldest-first and lead the merge's own records. + assert messages.index("[c1] c1 retrieval log") < messages.index( + "[c2] c2 retrieval log" + ) + # ...and they aren't left on the merged response, which would duplicate + # them once finish_query splices the log store in. + assert get_message_sync("rid").get("logs") == [] + + +def test_merge_messages_by_ids_direct_lookup_logs_not_left_on_message(mocker): + """The direct-lookup path returns the callback message as the accumulator + verbatim, so its logs have to be taken off it too.""" + query_graph = response_2["message"]["query_graph"] + _patch_sync_store(mocker) + from shepherd_utils.db import save_message_sync + + save_message_sync("qid", {"message": {"query_graph": copy.deepcopy(query_graph)}}) + save_message_sync("rid", copy.deepcopy(response_2)) + save_message_sync( + "c1", _callback_with_logs([{"level": "INFO", "message": "lookup log"}]) + ) + + _, log_entries = merge_messages_by_ids( + "test_ara", "qid", "rid", ["c1"], logging.INFO + ) + + assert "[c1] lookup log" in [entry.get("message") for entry in log_entries] + assert get_message_sync("rid").get("logs") == [] + + def test_merge_messages_by_id_delegates(mocker): """The single-callback entry point still works via the batched path.""" query_graph = response_1["message"]["query_graph"] diff --git a/tests/unit/test_shared_utils.py b/tests/unit/test_shared_utils.py index 9f8766d..fb14202 100644 --- a/tests/unit/test_shared_utils.py +++ b/tests/unit/test_shared_utils.py @@ -655,3 +655,39 @@ async def test_run_task_lifecycle_cancellation_does_not_route_failure(mocker): assert not mock_fail.called assert limiter.released + + +@pytest.mark.asyncio +async def test_tasks_for_one_query_do_not_duplicate_each_others_logs(redis_mock): + """A worker gets one task per callback for the same query, and every one of + them builds "its" logger from the same process-wide name. Each task's flush + must store only what that task logged: stacking a handler per task, and + flushing without emptying it, wrote the earlier tasks' records again every + time -- so a single retrieval log line showed up once per later callback. + """ + from shepherd_utils.db import get_logs, save_logs + from shepherd_utils.shared import _build_task_context + + def task(msg_id): + return (msg_id, {"query_id": "q1", "response_id": "rid", "log_level": "20"}) + + _, first_logger = _build_task_context("merge_message", "c", task("1-1"), 20) + first_logger.info("callback one retrieval log") + await save_logs("rid", first_logger) + + _, second_logger = _build_task_context("merge_message", "c", task("1-2"), 20) + assert second_logger is first_logger + # ...and the second task reused the handler rather than stacking another on + # top, which would queue every subsequent record once per attached handler. + assert ( + len([h for h in second_logger.handlers if h.name == "query_log_handler"]) == 1 + ) + second_logger.info("callback two retrieval log") + await save_logs("rid", second_logger) + + logs = await get_logs("rid", first_logger) + assert [entry["message"] for entry in logs] == [ + "callback one retrieval log", + "callback two retrieval log", + ] + first_logger.handlers.clear() diff --git a/workers/aragorn_lookup/worker.py b/workers/aragorn_lookup/worker.py index 50beea9..b23d588 100644 --- a/workers/aragorn_lookup/worker.py +++ b/workers/aragorn_lookup/worker.py @@ -103,8 +103,11 @@ async def run_async_lookup( message["callback"] = f"{settings.callback_host}/aragorn/callback/{callback_id}" - logger.debug( - f"""Sending lookup query to {settings.kg_retrieval_url} with callback {message['callback']}""" + # INFO, and tagged with the callback id: this is the head of the trail + # for one retrieval, and the callback handler, the retrieval's own logs + # and the merge all carry the same tag on the way back. + logger.info( + f"[{callback_id}] Sending lookup query to {settings.kg_retrieval_url}" ) try: response = await client.post( @@ -160,7 +163,9 @@ async def aragorn_lookup(task, logger: logging.Logger): # with open("./debug/direct_query.json", "w", encoding="utf-8") as f: # json.dump(message, f, indent=2) - logger.debug(f"""Sending lookup query to {settings.kg_retrieval_url}.""") + logger.info( + f"[{callback_id}] Sending lookup query to {settings.kg_retrieval_url}" + ) with tracer.start_as_current_span("aragorn.lookup") as span: span.set_attribute("callback_id", callback_id) async with httpx.AsyncClient(timeout=100) as client: @@ -193,7 +198,8 @@ async def aragorn_lookup(task, logger: logging.Logger): elif isinstance(response, AsyncResponse): if not response.success: logger.error( - f"Failed to do lookup, removing callback id: {response.error}" + f"[{response.callback_id}] Failed to do lookup, " + f"removing callback id: {response.error}" ) await remove_callback_id(response.callback_id, logger) else: diff --git a/workers/aragorn_pathfinder/worker.py b/workers/aragorn_pathfinder/worker.py index dc003d5..38d3a7e 100644 --- a/workers/aragorn_pathfinder/worker.py +++ b/workers/aragorn_pathfinder/worker.py @@ -215,7 +215,9 @@ async def shadowfax(task, logger: logging.Logger) -> str: retriever_query["callback"] = ( f"{settings.callback_host}/aragorn/callback/{callback_id}" ) - logger.debug(f"""Sending pathfinder query to {settings.kg_retrieval_url}.""") + logger.info( + f"[{callback_id}] Sending pathfinder query to {settings.kg_retrieval_url}" + ) with tracer.start_as_current_span("aragorn.pathfinder") as span: span.set_attribute("callback_id", callback_id) async with httpx.AsyncClient(timeout=100) as client: diff --git a/workers/bte_lookup/worker.py b/workers/bte_lookup/worker.py index 5e3c1e4..8e8fc85 100644 --- a/workers/bte_lookup/worker.py +++ b/workers/bte_lookup/worker.py @@ -103,8 +103,11 @@ async def run_async_lookup( message["callback"] = f"{settings.callback_host}/bte/callback/{callback_id}" - logger.debug( - f"""Sending lookup query to {settings.kg_retrieval_url} with callback {message['callback']}""" + # INFO, and tagged with the callback id: this is the head of the trail + # for one retrieval, and the callback handler, the retrieval's own logs + # and the merge all carry the same tag on the way back. + logger.info( + f"[{callback_id}] Sending lookup query to {settings.kg_retrieval_url}" ) try: response = await client.post( @@ -153,6 +156,9 @@ async def bte_lookup(task, logger: logging.Logger): await add_callback_id(query_id, callback_id, otel, logger) message["callback"] = f"{settings.callback_host}/bte/callback/{callback_id}" + logger.info( + f"[{callback_id}] Sending lookup query to {settings.kg_retrieval_url}" + ) with tracer.start_as_current_span("bte.lookup") as span: span.set_attribute("callback_id", callback_id) async with httpx.AsyncClient(timeout=100) as client: @@ -185,7 +191,8 @@ async def bte_lookup(task, logger: logging.Logger): elif isinstance(response, AsyncResponse): if not response.success: logger.error( - f"Failed to do lookup, removing callback id: {response.error}" + f"[{response.callback_id}] Failed to do lookup, " + f"removing callback id: {response.error}" ) await remove_callback_id(response.callback_id, logger) else: diff --git a/workers/merge_message/worker.py b/workers/merge_message/worker.py index bdfa0ba..9986883 100644 --- a/workers/merge_message/worker.py +++ b/workers/merge_message/worker.py @@ -30,7 +30,7 @@ save_logs, save_message_sync, ) -from shepherd_utils.logger import QueryLogger, get_worker_logger +from shepherd_utils.logger import QueryLogger, get_query_handler, get_worker_logger from shepherd_utils.otel import setup_tracer from shepherd_utils.process_pool import ProcessPoolManager from shepherd_utils.shared import filter_kgraph_orphans, get_tasks, merge_kgraph @@ -651,6 +651,79 @@ def merge_messages( raise TypeError("Unsupported query type.") +# TRAPI LogEntry level names, mapped to the ``logging`` levels they filter +# against. Anything else (a missing or unrecognized level) isn't filtered. +TRAPI_LOG_LEVELS = { + "DEBUG": logging.DEBUG, + "INFO": logging.INFO, + "WARNING": logging.WARNING, + "ERROR": logging.ERROR, + "CRITICAL": logging.CRITICAL, +} + + +def take_callback_logs( + callback_response: dict[str, Any], + callback_id: str, + log_level: int, + logger: logging.Logger, +) -> list[dict]: + """Take the log entries a subservice returned with its callback message. + + The KG retrieval services report what they did -- which KPs they called, + what timed out, why an edge came back empty -- in the TRAPI ``logs`` list of + the response they post back to ``/callback``. Merging only ever builds a + fresh message, so those entries used to be dropped on the floor and never + reached the query's log list. Lift them out here so the caller can fold them + in alongside the merge's own records. + + Each entry is tagged with the callback id it arrived under, the same + ``[callback_id]`` prefix the lookup worker and the callback handler use. + One query fans out into many retrievals whose logs all land in one list, so + the tag is what ties a line back to the retrieval that emitted it -- and + joins it to the dispatch and merge lines on either side. + + The field is blanked on the callback message itself: the merged response's + logs are spliced in from the log store when the query finishes, so a + leftover list here would show up twice in that payload (the direct-lookup + path returns the callback message as the accumulator verbatim). + + Entries below the query's requested ``log_level`` are dropped, matching what + the rest of the pipeline persists. + """ + logs = callback_response.pop("logs", None) + callback_response["logs"] = [] + if not logs: + return [] + if not isinstance(logs, list): + logger.warning( + f"Callback {callback_id} returned a non-list 'logs' field; ignoring it." + ) + return [] + entries = [] + for log in logs: + if not isinstance(log, dict): + # Not a TRAPI LogEntry, but there's still something a human wrote in + # there -- keep it rather than silently dropping it. + entries.append({"message": f"[{callback_id}] {log}", "level": "INFO"}) + continue + level = TRAPI_LOG_LEVELS.get(str(log.get("level", "")).upper()) + if level is not None and level < log_level: + continue + message = log.get("message") + log["message"] = f"[{callback_id}] {message}" if message else f"[{callback_id}]" + entries.append(log) + max_logs = settings.merge_max_callback_logs + if max_logs > 0 and len(entries) > max_logs: + dropped = len(entries) - max_logs + entries = entries[:max_logs] + logger.warning( + f"Callback {callback_id} returned {dropped} log entries beyond the " + f"{max_logs}-entry cap; the rest were dropped." + ) + return entries + + # --------------------------------------------------------------------------- # Worker entry point # @@ -680,11 +753,13 @@ def merge_messages_by_ids( callback ids that were actually folded in, so the caller knows which to clear from the ready index and the callbacks table (a single missing callback is skipped rather than aborting the whole batch). ``log_entries`` - is the list of formatted log records produced during the merge, oldest - first -- this runs in a ProcessPoolExecutor child, so its logger can't be + is the list of log records to add to the query's log list, oldest first: + the ones each callback carried back from the KG retrieval that produced it + (see ``take_callback_logs``), followed by the ones this merge emitted + itself. This runs in a ProcessPoolExecutor child, so its logger can't be the parent's query logger; instead we attach a fresh ``QueryLogHandler`` here and hand its contents back across the process boundary for the parent - to fold into the query's log list. + to fold in. """ # A logger.getLogger call in a child returns the same object for the whole # process life, so attach a call-scoped handler and remove it in finally -- @@ -704,6 +779,7 @@ def merge_messages_by_ids( original_query_graph = original_query["message"]["query_graph"] merged: list[str] = [] + callback_log_entries: list[dict] = [] for callback_id in callback_ids: try: callback_response = get_message_sync(callback_id) @@ -712,6 +788,11 @@ def merge_messages_by_ids( f"Missing callback {callback_id} while folding; skipping." ) continue + callback_log_entries.extend( + take_callback_logs( + callback_response, callback_id, log_level, worker_logger + ) + ) accumulator = merge_messages( target, original_query_graph, @@ -723,12 +804,10 @@ def merge_messages_by_ids( if merged: save_message_sync(response_id, accumulator) - # contents() is newest-first (emit appendlefts); hand back oldest-first - # so the parent can appendleft them in order and keep the queue's - # newest-first invariant. - log_entries = list(query_log_handler.contents()) - log_entries.reverse() - return merged, log_entries + # drain() hands back oldest-first, which is the order the parent's + # handler.ingest wants. The retrieval logs describe work that happened + # before this merge, so they lead. + return merged, callback_log_entries + query_log_handler.drain() finally: worker_logger.removeHandler(query_log_handler) @@ -800,14 +879,7 @@ def _ingest_merge_logs(logger, entries): """ if not entries: return - handler = next( - ( - h - for h in logger.handlers - if getattr(h, "name", None) == "query_log_handler" - ), - None, - ) + handler = get_query_handler(logger) if handler is not None: handler.ingest(entries)