278 lines
8.9 KiB
Python
278 lines
8.9 KiB
Python
import json
|
|
import logging
|
|
import traceback
|
|
from typing import Annotated, Literal
|
|
|
|
from pydantic import BaseModel, ConfigDict, Field, ValidationError, model_validator
|
|
|
|
from agent import Agent, AgentMessage
|
|
from paths import debug_path, prompt_path
|
|
from utils import Event, Timeline
|
|
from windowizer import Window, Windowizer
|
|
|
|
|
|
class _StrictModel(BaseModel):
|
|
"""Base model for schemas returned by the agent."""
|
|
|
|
model_config = ConfigDict(extra="forbid", strict=True)
|
|
|
|
|
|
class HeadingElement(_StrictModel):
|
|
"""A document section heading."""
|
|
|
|
type: Literal["heading"]
|
|
level: Literal[2, 3, 4]
|
|
text: str = Field(min_length=1)
|
|
|
|
|
|
class ParagraphElement(_StrictModel):
|
|
"""A paragraph containing one complete thought."""
|
|
|
|
type: Literal["paragraph"]
|
|
text: str = Field(min_length=1)
|
|
|
|
|
|
class UnorderedListElement(_StrictModel):
|
|
"""A list whose item order is not significant."""
|
|
|
|
type: Literal["unordered"]
|
|
items: list[str] = Field(min_length=1)
|
|
|
|
@model_validator(mode="after")
|
|
def validate_items(self) -> "UnorderedListElement":
|
|
if any(not item.strip() for item in self.items):
|
|
raise ValueError("List items must not be empty")
|
|
return self
|
|
|
|
|
|
class OrderedListElement(_StrictModel):
|
|
"""A list whose item order is significant."""
|
|
|
|
type: Literal["ordered"]
|
|
items: list[str] = Field(min_length=1)
|
|
|
|
@model_validator(mode="after")
|
|
def validate_items(self) -> "OrderedListElement":
|
|
if any(not item.strip() for item in self.items):
|
|
raise ValueError("List items must not be empty")
|
|
return self
|
|
|
|
|
|
class DefinitionElement(_StrictModel):
|
|
"""A definition of a term or concept."""
|
|
|
|
type: Literal["definition"]
|
|
term: str = Field(min_length=1)
|
|
text: str = Field(min_length=1)
|
|
|
|
|
|
class ImportantElement(_StrictModel):
|
|
"""An important statement that should stand out in the document."""
|
|
|
|
type: Literal["important"]
|
|
text: str = Field(min_length=1)
|
|
|
|
|
|
class ImageElement(_StrictModel):
|
|
"""A reference to an image stored in a visual timeline event."""
|
|
|
|
type: Literal["image"]
|
|
event_id: int
|
|
|
|
|
|
StructureElement = (
|
|
HeadingElement
|
|
| ParagraphElement
|
|
| UnorderedListElement
|
|
| OrderedListElement
|
|
| DefinitionElement
|
|
| ImportantElement
|
|
| ImageElement
|
|
)
|
|
|
|
|
|
class Structure(_StrictModel):
|
|
"""Structure of the future Markdown document."""
|
|
|
|
elements: list[StructureElement]
|
|
|
|
|
|
class _Pending(_StrictModel):
|
|
"""An unfinished thought carried over to the next window."""
|
|
|
|
event_ids: list[int] = Field(min_length=1)
|
|
summary: str = Field(min_length=1)
|
|
|
|
@model_validator(mode="after")
|
|
def validate_event_ids(self) -> "_Pending":
|
|
if len(self.event_ids) != len(set(self.event_ids)):
|
|
raise ValueError("pending.event_ids must not contain duplicates")
|
|
return self
|
|
|
|
|
|
class _BuilderContext(_StrictModel):
|
|
"""Small amount of document state preserved between agent calls."""
|
|
|
|
current_section: str | None = Field(default=None, max_length=120)
|
|
current_subsection: str | None = Field(default=None, max_length=120)
|
|
recent_headings: list[str] = Field(default_factory=list, max_length=6)
|
|
pending: _Pending | None = None
|
|
|
|
@model_validator(mode="after")
|
|
def validate_recent_headings(self) -> "_BuilderContext":
|
|
if any(not heading.strip() or len(heading) > 120 for heading in self.recent_headings):
|
|
raise ValueError("recent_headings must contain non-empty strings up to 120 characters")
|
|
return self
|
|
|
|
|
|
class _BuildResult(_StrictModel):
|
|
"""Result returned by the agent for a single window."""
|
|
|
|
elements: list[StructureElement]
|
|
context: _BuilderContext
|
|
|
|
|
|
class StructureBuilder:
|
|
"""Build a document structure from a timeline of lecture events."""
|
|
|
|
DEBUG_ID = 0
|
|
MAX_RETRIES = 5
|
|
|
|
def __init__(self, agent: Agent, windowizer: Windowizer[Event]) -> None:
|
|
"""Create a structure builder.
|
|
|
|
Args:
|
|
agent: Agent used to transform timeline windows.
|
|
windowizer: Windowizer used to split timeline events.
|
|
"""
|
|
self._agent = agent
|
|
self._windowizer = windowizer
|
|
self._system_prompt = AgentMessage(
|
|
content=prompt_path("structure_builder.md").read_text(encoding="utf-8"),
|
|
role="system",
|
|
)
|
|
|
|
@staticmethod
|
|
def _validate_result(result: _BuildResult, window: Window[Event]) -> None:
|
|
"""Validate constraints that depend on the current input window."""
|
|
usable_visual_ids = {
|
|
event.id
|
|
for event in window.past + window.present
|
|
if event.type == "vis"
|
|
}
|
|
for element in result.elements:
|
|
if isinstance(element, ImageElement) and element.event_id not in usable_visual_ids:
|
|
raise ValueError(
|
|
"image.event_id must reference a vis event from past or present"
|
|
)
|
|
|
|
if result.context.pending is None:
|
|
return
|
|
|
|
window_ids = {
|
|
event.id
|
|
for event in window.past + window.present + window.future
|
|
}
|
|
invalid_pending_ids = (
|
|
set(result.context.pending.event_ids) - window_ids
|
|
)
|
|
if invalid_pending_ids:
|
|
raise ValueError(
|
|
"pending.event_ids must reference events from the current window"
|
|
)
|
|
|
|
@staticmethod
|
|
def _keep_window_references(
|
|
result: _BuildResult,
|
|
window: Window[Event],
|
|
) -> _BuildResult:
|
|
"""Drop impossible image references and repair technical pending IDs."""
|
|
visual_ids = {
|
|
event.id
|
|
for event in window.past + window.present
|
|
if event.type == "vis"
|
|
}
|
|
elements: list[StructureElement] = []
|
|
for element in result.elements:
|
|
if isinstance(element, ImageElement) and element.event_id not in visual_ids:
|
|
logging.warning("Ignoring image outside the current window: id=%s", element.event_id)
|
|
continue
|
|
elements.append(element)
|
|
|
|
context = result.context
|
|
if context.pending is not None:
|
|
window_ids = {
|
|
event.id
|
|
for event in window.past + window.present + window.future
|
|
}
|
|
pending_ids = [
|
|
event_id
|
|
for event_id in context.pending.event_ids
|
|
if event_id in window_ids
|
|
]
|
|
if pending_ids != context.pending.event_ids:
|
|
logging.warning(
|
|
"Trimming pending IDs to the current window: %s -> %s",
|
|
context.pending.event_ids,
|
|
pending_ids,
|
|
)
|
|
pending = (
|
|
context.pending.model_copy(update={"event_ids": pending_ids})
|
|
if pending_ids else None
|
|
)
|
|
context = context.model_copy(update={"pending": pending})
|
|
|
|
return result.model_copy(update={"elements": elements, "context": context})
|
|
|
|
def _build_window(self, window: Window[Event]) -> _BuildResult:
|
|
"""Build document elements from a single timeline window."""
|
|
messages = [
|
|
self._system_prompt,
|
|
AgentMessage(
|
|
content=json.dumps(
|
|
window.model_dump(mode="json"),
|
|
indent=2,
|
|
ensure_ascii=False,
|
|
),
|
|
role="user",
|
|
),
|
|
]
|
|
|
|
debug_dir = debug_path("StructureBuilder", StructureBuilder.DEBUG_ID)
|
|
if debug_dir:
|
|
StructureBuilder.DEBUG_ID += 1
|
|
(debug_dir / "request.txt").write_text(messages[1].content, encoding="utf-8")
|
|
|
|
retries_left = self.MAX_RETRIES
|
|
while retries_left > 0:
|
|
retries_left -= 1
|
|
response = self._agent.completion(messages=messages, response_format=_BuildResult)
|
|
if debug_dir:
|
|
(debug_dir / f"{retries_left}-retries-left.txt").write_text(
|
|
response,
|
|
encoding="utf-8",
|
|
)
|
|
|
|
try:
|
|
result = _BuildResult.model_validate_json(response)
|
|
result = self._keep_window_references(result, window)
|
|
self._validate_result(result, window)
|
|
return result
|
|
except (ValidationError, ValueError, TypeError):
|
|
traceback.print_exc()
|
|
|
|
raise RuntimeError("Agent has failed to provide valid schema too many times")
|
|
|
|
def build(self, timeline: Timeline) -> Structure:
|
|
"""Build the ordered structure of a future Markdown document."""
|
|
elements: list[StructureElement] = []
|
|
context = _BuilderContext()
|
|
|
|
for window in self._windowizer.windowize(timeline.events):
|
|
window.context = context.model_dump(mode="json")
|
|
intermediate = self._build_window(window)
|
|
elements.extend(intermediate.elements)
|
|
context = intermediate.context
|
|
|
|
return Structure(elements=elements)
|