Skip to content

Commit 3422dc1

Browse files
committed
minor fixes
1 parent 9543006 commit 3422dc1

5 files changed

Lines changed: 60 additions & 43 deletions

File tree

‎src/inspect_ai/_control/eval_state.py‎

Lines changed: 16 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -216,6 +216,16 @@ class EvalState:
216216
total_messages: int = 0
217217
"""Cumulative message count, accumulated like :attr:`total_tokens`."""
218218

219+
@property
220+
def terminal(self) -> int:
221+
"""Samples that reached a terminal outcome (completed/errored/cancelled).
222+
223+
The single definition of the terminal sum — used by both
224+
:attr:`is_finished` and :func:`finalize_eval` so a future bucket
225+
can't be added to one and missed by the other.
226+
"""
227+
return self.completed + self.errored + self.cancelled
228+
219229
@property
220230
def is_finished(self) -> bool:
221231
"""True once every sample has terminated (success, error, or cancel).
@@ -226,7 +236,7 @@ def is_finished(self) -> bool:
226236
early-finish risk for ``total > 0`` — ``0 >= total`` is already False
227237
until enough samples terminate.
228238
"""
229-
return self.completed + self.errored + self.cancelled >= self.total
239+
return self.terminal >= self.total
230240

231241

232242
# Module-level registry. Keyed by eval_id. A process can host multiple
@@ -477,9 +487,7 @@ def finalize_eval(eval_id: str) -> None:
477487
with _lock:
478488
state = _eval_states.get(eval_id)
479489
if state is not None:
480-
shortfall = state.total - (
481-
state.completed + state.errored + state.cancelled
482-
)
490+
shortfall = state.total - state.terminal
483491
if shortfall > 0:
484492
state.cancelled += shortfall
485493
_maybe_mark_finished(state)
@@ -488,9 +496,10 @@ def finalize_eval(eval_id: str) -> None:
488496
def _maybe_mark_finished(state: EvalState) -> None:
489497
"""Stamp ``completed_at`` when every sample has terminated.
490498
491-
Fires the first time ``completed + errored`` reaches ``total``;
492-
later updates are no-ops so a late counter update from a teardown
493-
race doesn't overwrite the original finish time. Also drops
499+
Fires the first time the terminal sum (``completed + errored +
500+
cancelled``) reaches ``total``; later updates are no-ops so a late
501+
counter update from a teardown race doesn't overwrite the original
502+
finish time. Also drops
494503
``sample_ids`` — a finished eval has no pending samples, so the
495504
planned-id list is dead weight (it's retained on the state until the
496505
run boundary clears it). Caller must hold the registry lock.

‎src/inspect_ai/_control/events.py‎

Lines changed: 4 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -2,7 +2,10 @@
22
33
Backs ``GET /evals/<id>/sample/events`` (and ``inspect ctl events``): a
44
**cursored-pull** window over one sample's events, read from its live
5-
``Transcript`` while running and from the on-disk log once terminal.
5+
``Transcript`` while running, and once terminal from the recorder's
6+
sample, the realtime buffer (via the eval's events provider — the
7+
streaming-completion path retains an event-less recorder sample), or the
8+
on-disk log (see ``_logged_source``).
69
710
The cursor is an opaque token = ``(source nonce, absolute event offset)``.
811
The offset indexes the *unfiltered* event sequence; type / time filters are

‎src/inspect_ai/_control/state.py‎

Lines changed: 7 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -37,6 +37,7 @@
3737
from functools import partial
3838
from typing import TYPE_CHECKING, Any
3939

40+
from inspect_ai._util._async import tg_collect
4041
from inspect_ai._util.error import is_cancellation_message
4142
from inspect_ai._util.file import local_path
4243

@@ -87,8 +88,6 @@ async def current_eval_summaries(started_at: float) -> list[dict[str, Any]]:
8788
# governed by the filesystem's connection pool.
8889
deferred = [s for s in states if s.deferred_sample_stats is not None]
8990
if deferred:
90-
from inspect_ai._util._async import tg_collect
91-
9291
await tg_collect(
9392
[partial(resolve_deferred_sample_stats, state) for state in deferred]
9493
)
@@ -117,8 +116,8 @@ async def current_eval_summaries(started_at: float) -> list[dict[str, Any]]:
117116
# eval_ids covered by some grouped state — used to attribute live
118117
# samples to their group when building the per-group summary.
119118
eval_id_to_group: dict[str, str] = {}
120-
for group_key, states in states_by_group.items():
121-
for state in states:
119+
for group_key, group_states in states_by_group.items():
120+
for state in group_states:
122121
eval_id_to_group[state.eval_id] = group_key
123122

124123
# Live samples whose eval has no registered EvalState (eg. a brand-
@@ -128,20 +127,20 @@ async def current_eval_summaries(started_at: float) -> list[dict[str, Any]]:
128127

129128
summaries: list[dict[str, Any]] = []
130129

131-
for group_key, states in states_by_group.items():
130+
for group_key, group_states in states_by_group.items():
132131
# Latest attempt = last registered. Retries are sequential (a
133132
# retry registers only after the prior attempt finishes), and
134133
# `get_eval_states()` preserves registration order, so the tail
135134
# is the current attempt. Selecting by `completed_at` would wrongly
136135
# prefer a finished earlier attempt over a still-running retry
137136
# (whose `completed_at` is None).
138-
latest = states[-1]
139-
attempts = len(states)
137+
latest = group_states[-1]
138+
attempts = len(group_states)
140139

141140
# Live samples: pull from every attempt in the group (only the
142141
# latest will normally have any, but be defensive).
143142
group_samples: list[ActiveSample] = []
144-
for state in states:
143+
for state in group_states:
145144
group_samples.extend(samples_by_eval.get(state.eval_id, []))
146145

147146
summaries.append(

‎src/inspect_ai/_eval/evalset.py‎

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -25,9 +25,6 @@
2525
clear_all_eval_states,
2626
register_completed_eval,
2727
)
28-
29-
if TYPE_CHECKING:
30-
from inspect_ai._control.eval_state import DeferredStatsProvider
3128
from inspect_ai._control.server import (
3229
control_server,
3330
release_requested,
@@ -95,6 +92,9 @@
9592
from .task.task import PreviousTask, resolve_epochs
9693
from .task.tasks import Tasks
9794

95+
if TYPE_CHECKING:
96+
from inspect_ai._control.eval_state import DeferredStatsProvider
97+
9898
logger = logging.getLogger(__name__)
9999

100100

‎src/inspect_ai/_eval/task/run.py‎

Lines changed: 30 additions & 24 deletions
Original file line numberDiff line numberDiff line change
@@ -356,6 +356,32 @@ async def task_run(options: TaskRunOptions, task_cancel: TaskCancel | None) -> E
356356
# deleted below — used by register_eval and carry_forward_unlogged_samples
357357
sample_ids = [s.id for s in dataset if s.id is not None]
358358

359+
async def finish_task_log(
360+
status: EvalStatus,
361+
stats: EvalStats,
362+
results: EvalResults | None = None,
363+
reductions: list[EvalSampleReductions] | None = None,
364+
error: EvalError | None = None,
365+
) -> EvalLog:
366+
"""Finish via ``_finish_task_log`` with the run-invariant context.
367+
368+
Bound once here so every terminal branch finishes with the same
369+
logger / sample-source / planned-ids context — a new branch can't
370+
accidentally thread a stale or divergent value.
371+
"""
372+
return await _finish_task_log(
373+
logger=logger,
374+
sample_source=options.sample_source,
375+
sample_ids=sample_ids,
376+
epochs=epochs,
377+
log_images=log_images,
378+
status=status,
379+
stats=stats,
380+
results=results,
381+
reductions=reductions,
382+
error=error,
383+
)
384+
359385
# handle sample errors (raise as required). use total_samples (sliced
360386
# dataset * epochs) as the denominator for fractional fail_on_error so
361387
# the mid-run abort threshold matches the end-of-run check below.
@@ -743,12 +769,7 @@ async def create_sample_state(
743769
)
744770

745771
# finish
746-
eval_log = await _finish_task_log(
747-
logger=logger,
748-
sample_source=options.sample_source,
749-
sample_ids=sample_ids,
750-
epochs=epochs,
751-
log_images=log_images,
772+
eval_log = await finish_task_log(
752773
status="error" if mark_log_as_error else "success",
753774
stats=stats,
754775
results=results,
@@ -789,12 +810,7 @@ async def create_sample_state(
789810
f"Task cancelled by user ({task_cancel.cancel_type})"
790811
)
791812
error = eval_error(cancel_ex, TerminateTaskError, cancel_ex, None)
792-
eval_log = await _finish_task_log(
793-
logger=logger,
794-
sample_source=options.sample_source,
795-
sample_ids=sample_ids,
796-
epochs=epochs,
797-
log_images=log_images,
813+
eval_log = await finish_task_log(
798814
status="error",
799815
stats=stats,
800816
results=results,
@@ -811,12 +827,7 @@ async def create_sample_state(
811827
)
812828
else:
813829
# External cancellation (ctrl+c)
814-
eval_log = await _finish_task_log(
815-
logger=logger,
816-
sample_source=options.sample_source,
817-
sample_ids=sample_ids,
818-
epochs=epochs,
819-
log_images=log_images,
830+
eval_log = await finish_task_log(
820831
status="cancelled",
821832
stats=stats,
822833
results=results,
@@ -840,12 +851,7 @@ async def create_sample_state(
840851
collect_eval_data(stats)
841852

842853
# finish with error status
843-
eval_log = await _finish_task_log(
844-
logger=logger,
845-
sample_source=options.sample_source,
846-
sample_ids=sample_ids,
847-
epochs=epochs,
848-
log_images=log_images,
854+
eval_log = await finish_task_log(
849855
status="error",
850856
stats=stats,
851857
results=results,

0 commit comments

Comments
 (0)