From 67fddeaaf4ccfbabfc186fd88acc69ae276ff4b1 Mon Sep 17 00:00:00 2001 From: vegu-ai-tools <152010387+vegu-ai-tools@users.noreply.github.com> Date: Mon, 1 Jun 2026 13:26:42 +0300 Subject: [PATCH] Refactor world state management by introducing mixins for reinforcements and pin conditions --- src/talemate/agents/world_state/__init__.py | 368 +----------------- .../agents/world_state/pin_conditions.py | 224 +++++++++++ .../agents/world_state/reinforcements.py | 197 ++++++++++ 3 files changed, 429 insertions(+), 360 deletions(-) create mode 100644 src/talemate/agents/world_state/pin_conditions.py create mode 100644 src/talemate/agents/world_state/reinforcements.py diff --git a/src/talemate/agents/world_state/__init__.py b/src/talemate/agents/world_state/__init__.py index f4fd2bab..a99b7f89 100644 --- a/src/talemate/agents/world_state/__init__.py +++ b/src/talemate/agents/world_state/__init__.py @@ -1,30 +1,21 @@ from __future__ import annotations -import json -from typing import TYPE_CHECKING - import isodate import structlog import talemate.emit.async_signals import talemate.util as util from talemate.emit import emit -from talemate.events import GameLoopEvent from talemate.instance import get_agent from talemate.client import ClientBase from talemate.prompts import Prompt -from talemate.prompts.response import AnchorExtractor, ResponseSpec -from talemate.scene_message import ( - ReinforcementMessage, - TimePassageMessage, -) +from talemate.scene_message import TimePassageMessage from talemate.util.response import extract_list from talemate.agents.base import ( Agent, AgentAction, - AgentActionConfig, AgentEmission, DynamicInstruction, optimize_prompt_caching_action, @@ -38,12 +29,11 @@ from .character_progression import CharacterProgressionMixin from .avatars import AvatarMixin from .snapshot import WorldStateSnapshotMixin from .snapshot import SNAPSHOT_TASK as SNAPSHOT_TASK +from .reinforcements import WorldStateReinforcementsMixin +from .pin_conditions import WorldStatePinConditionsMixin from .websocket_handler import WorldStateWebsocketHandler import talemate.agents.world_state.nodes -if TYPE_CHECKING: - from talemate.tale_mate import Character - log = structlog.get_logger("talemate.agents.world_state") talemate.emit.async_signals.register("agent.world_state.time") @@ -71,6 +61,8 @@ class TimePassageEmission(WorldStateAgentEmission): class WorldStateAgent( MemoryRAGMixin, WorldStateSnapshotMixin, + WorldStateReinforcementsMixin, + WorldStatePinConditionsMixin, CharacterProgressionMixin, AvatarMixin, Agent, @@ -88,41 +80,10 @@ class WorldStateAgent( actions = { "prompt_caching": optimize_prompt_caching_action(), } - # Inserted before the other actions to keep update_world_state right - # after prompt_caching in the UI. + # add_actions calls are ordered to match the action list shown in the UI. WorldStateSnapshotMixin.add_actions(actions) - actions.update( - { - "update_reinforcements": AgentAction( - enabled=True, - can_be_disabled=True, - container=True, - icon="mdi-image-auto-adjust", - label="Update state reinforcements", - description="Will attempt to update any due state reinforcements.", - config={}, - ), - "check_pin_conditions": AgentAction( - enabled=True, - can_be_disabled=True, - container=True, - icon="mdi-pin", - label="Update conditional context pins", - description="Will evaluate context pins conditions and toggle those pins accordingly. Runs automatically every N turns.", - config={ - "turns": AgentActionConfig( - type="number", - label="Turns", - description="Number of turns to wait before checking conditions.", - value=2, - min=1, - max=100, - step=1, - ) - }, - ), - } - ) + WorldStateReinforcementsMixin.add_actions(actions) + WorldStatePinConditionsMixin.add_actions(actions) MemoryRAGMixin.add_actions(actions) CharacterProgressionMixin.add_actions(actions) AvatarMixin.add_actions(actions) @@ -147,22 +108,6 @@ class WorldStateAgent( def experimental(self): return True - @property - def update_reinforcements_enabled(self) -> bool: - return self.resolve_enabled("update_reinforcements") - - @property - def check_pin_conditions_enabled(self) -> bool: - return self.resolve_enabled("check_pin_conditions") - - @property - def check_pin_conditions_turns(self): - return self.resolve_config("check_pin_conditions", "turns") - - def connect(self, scene): - super().connect(scene) - talemate.emit.async_signals.get("game_loop").connect(self.on_game_loop) - async def advance_time( self, duration: str, narrative: str = None ) -> TimePassageMessage: @@ -191,44 +136,6 @@ class WorldStateAgent( return message - async def on_game_loop(self, emission: GameLoopEvent): - """ - Called once per scene-loop round. - """ - - if not self.enabled: - return - - await self.auto_update_reinforcments() - await self.auto_check_pin_conditions() - - async def auto_update_reinforcments(self): - if not self.enabled: - return - - if not self.update_reinforcements_enabled: - return - - await self.update_reinforcements() - - async def auto_check_pin_conditions(self): - if not self.enabled: - return - - if not self.check_pin_conditions_enabled: - return - - if ( - self.next_pin_check % self.check_pin_conditions_turns != 0 - or self.next_pin_check == 0 - ): - self.next_pin_check += 1 - return - - self.next_pin_check = 0 - - await self.check_pin_conditions() - @set_processing async def analyze_text_and_extract_context( self, @@ -485,265 +392,6 @@ class WorldStateAgent( extracted["response"], max_attributes=max_attributes ) - @set_processing - async def update_reinforcements(self, force: bool = False, reset: bool = False): - """ - Queries due worldstate re-inforcements - """ - - for reinforcement in self.scene.world_state.reinforce: - # Skip character reinforcements if require_active is True and character is not active - if ( - reinforcement.require_active - and reinforcement.character - and not self.scene.character_is_active(reinforcement.character) - ): - continue - - if reinforcement.due <= 0 or force: - await self.update_reinforcement( - reinforcement.question, reinforcement.character, reset=reset - ) - else: - reinforcement.due -= 1 - - @set_processing - async def update_reinforcement( - self, question: str, character: "str | Character" = None, reset: bool = False - ) -> str: - """ - Queries a single re-inforcement - """ - - if isinstance(character, self.scene.Character): - character = character.name - - message = None - idx, reinforcement = await self.scene.world_state.find_reinforcement( - question, character - ) - - if not reinforcement: - log.warning( - "Reinforcement not found", question=question, character=character - ) - return - - message = ReinforcementMessage(message="") - message.set_source( - "world_state", - "update_reinforcement", - question=question, - character=character, - ) - - if reset and reinforcement.insert == "sequential": - self.scene.pop_history( - typ="reinforcement", meta_hash=message.meta_hash, all=True - ) - - if reinforcement.insert == "sequential": - kind = "analyze_freeform_medium_short" - else: - kind = "analyze_freeform" - - response, extracted = await Prompt.request( - "world_state.update-reinforcements", - self.client, - kind, - vars={ - "scene": self.scene, - "max_tokens": self.client.max_token_length, - "question": reinforcement.question, - "instructions": reinforcement.instructions or "", - "character": ( - self.scene.get_character(reinforcement.character) - if reinforcement.character - else None - ), - "answer": (reinforcement.answer if not reset else None) or "", - "reinforcement": reinforcement, - }, - response_spec=ResponseSpec( - extractors={ - "response": AnchorExtractor( - left="", - right="", - fallback_to_full=True, - ), - }, - ), - ) - - answer = extracted["response"] - - reinforcement.answer = answer - reinforcement.due = reinforcement.interval - - # remove any recent previous reinforcement message with same question - # to avoid overloading the near history with reinforcement messages - if not reset: - self.scene.pop_history( - typ="reinforcement", meta_hash=message.meta_hash, max_iterations=10 - ) - - if reinforcement.insert == "sequential": - # insert the reinforcement message at the current position - message.message = answer - log.debug("update_reinforcement", message=message, reset=reset) - await self.scene.push_history(message) - - # if reinforcement has a character name set, update the character detail - if reinforcement.character: - character = self.scene.get_character(reinforcement.character) - await character.set_detail(reinforcement.question, answer) - - else: - # set world entry - await self.scene.world_state_manager.save_world_entry( - reinforcement.question, - reinforcement.as_context_line, - {}, - ) - - self.scene.world_state.emit() - - return message - - @set_processing - async def check_pin_conditions( - self, - ): - """ - Checks if any context pin conditions - """ - - log.debug("check_pin_conditions", turns=self.check_pin_conditions_turns) - - world_state = self.scene.world_state - - state_change = False - - # Build list of pins to check, honoring decay semantics - pins_to_check = {} - for entry_id, pin in world_state.pins.items(): - # Skip game-state-controlled pins from the LLM loop - if pin.gamestate_condition: - continue - # Initialize countdown if active with decay but no due set - if pin.active and pin.decay and not pin.decay_due: - pin.decay_due = pin.decay - - # Only pins with conditions are checked by the LLM - if not pin.condition: - continue - - # If pin is active and has decay, skip checks until it's about to decay (decay_due == 1) - if ( - pin.active - and pin.decay - and (pin.decay_due is not None) - and pin.decay_due > 1 - ): - continue - - # Include pin for checking when it has no decay, is inactive, or is about to decay - if (not pin.decay) or (not pin.active) or (pin.decay_due == 1): - pins_to_check[entry_id] = { - "condition": pin.condition, - "state": pin.condition_state, - } - - # Early return if nothing to check, but still tick decay - if not pins_to_check: - for entry_id, pin in world_state.pins.items(): - # Game-state-controlled pins do not decay - if pin.gamestate_condition: - continue - if pin.active and pin.decay: - if not pin.decay_due: - pin.decay_due = pin.decay - pin.decay_due -= self.check_pin_conditions_turns - log.debug("applying pin decay", pin=pin, decay_due=pin.decay_due) - if pin.decay_due <= 0: - log.debug("pin decay expired", pin=pin, decay_due=pin.decay_due) - pin.active = False - pin.decay_due = None - state_change = True - if state_change: - await self.scene.load_active_pins() - self.scene.emit_status() - return - - first_entry_id = list(pins_to_check.keys())[0] - - _, answers = await Prompt.request( - "world_state.check-pin-conditions", - self.client, - "analyze", - vars={ - "scene": self.scene, - "max_tokens": self.client.max_token_length, - "previous_states": json.dumps(pins_to_check, indent=2), - "coercion": {first_entry_id: {"condition": ""}}, - }, - ) - - # Apply LLM results - for entry_id, answer in answers.items(): - if entry_id not in world_state.pins: - log.debug( - "check_pin_conditions", - entry_id=entry_id, - answer=answer, - msg="entry_id not found in world_state.pins (LLM failed to produce a clean response)", - ) - continue - - log.debug("check_pin_conditions", entry_id=entry_id, answer=answer) - state = answer.get("state") - pin = world_state.pins[entry_id] - if state is True or ( - isinstance(state, str) and state.lower() in ["true", "yes", "y"] - ): - prev_state = pin.condition_state - pin.condition_state = True - if not pin.active: - state_change = True - pin.active = True - # Refresh decay countdown when condition is true and pin stays/turns active - if pin.decay: - pin.decay_due = pin.decay - if prev_state != pin.condition_state: - state_change = True - else: - if pin.condition_state is not False or pin.active: - pin.condition_state = False - pin.active = False - # Clear countdown when deactivated - pin.decay_due = None - state_change = True - - # Tick decay counters for all active pins with decay - for entry_id, pin in world_state.pins.items(): - # Game-state-controlled pins do not decay - if pin.gamestate_condition: - continue - if pin.active and pin.decay: - if not pin.decay_due: - pin.decay_due = pin.decay - # Decrement once per check cycle - pin.decay_due -= 1 - if pin.decay_due <= 0: - # Auto-deactivate on expiry - pin.active = False - pin.decay_due = None - state_change = True - - if state_change: - await self.scene.load_active_pins() - self.scene.emit_status() - @set_processing async def summarize_and_pin(self, message_id: int, num_messages: int = 3) -> str: """ diff --git a/src/talemate/agents/world_state/pin_conditions.py b/src/talemate/agents/world_state/pin_conditions.py new file mode 100644 index 00000000..8c989459 --- /dev/null +++ b/src/talemate/agents/world_state/pin_conditions.py @@ -0,0 +1,224 @@ +from __future__ import annotations + +import json + +import structlog + +import talemate.emit.async_signals +from talemate.events import GameLoopEvent +from talemate.prompts import Prompt + +from talemate.agents.base import AgentAction, AgentActionConfig, set_processing + +log = structlog.get_logger("talemate.agents.world_state") + + +class WorldStatePinConditionsMixin: + """ + World-state manager agent mixin that handles conditional context pins — + evaluating each pin's condition via the LLM, toggling the pin accordingly, + and ticking the decay countdown that auto-deactivates stale pins. + """ + + @classmethod + def add_actions(cls, actions: dict[str, AgentAction]): + actions["check_pin_conditions"] = AgentAction( + enabled=True, + can_be_disabled=True, + container=True, + icon="mdi-pin", + label="Update conditional context pins", + description="Will evaluate context pins conditions and toggle those pins accordingly. Runs automatically every N turns.", + config={ + "turns": AgentActionConfig( + type="number", + label="Turns", + description="Number of turns to wait before checking conditions.", + value=2, + min=1, + max=100, + step=1, + ) + }, + ) + + # config property helpers + + @property + def check_pin_conditions_enabled(self) -> bool: + return self.resolve_enabled("check_pin_conditions") + + @property + def check_pin_conditions_turns(self): + return self.resolve_config("check_pin_conditions", "turns") + + # signal connect + + def connect(self, scene): + super().connect(scene) + talemate.emit.async_signals.get("game_loop").connect( + self.on_game_loop_check_pin_conditions + ) + + async def on_game_loop_check_pin_conditions(self, emission: GameLoopEvent): + """ + Called once per scene-loop round. + """ + if not self.enabled: + return + + await self.auto_check_pin_conditions() + + # methods + + async def auto_check_pin_conditions(self): + if not self.enabled: + return + + if not self.check_pin_conditions_enabled: + return + + if ( + self.next_pin_check % self.check_pin_conditions_turns != 0 + or self.next_pin_check == 0 + ): + self.next_pin_check += 1 + return + + self.next_pin_check = 0 + + await self.check_pin_conditions() + + @set_processing + async def check_pin_conditions( + self, + ): + """ + Checks if any context pin conditions + """ + + log.debug("check_pin_conditions", turns=self.check_pin_conditions_turns) + + world_state = self.scene.world_state + + state_change = False + + # Build list of pins to check, honoring decay semantics + pins_to_check = {} + for entry_id, pin in world_state.pins.items(): + # Skip game-state-controlled pins from the LLM loop + if pin.gamestate_condition: + continue + # Initialize countdown if active with decay but no due set + if pin.active and pin.decay and not pin.decay_due: + pin.decay_due = pin.decay + + # Only pins with conditions are checked by the LLM + if not pin.condition: + continue + + # If pin is active and has decay, skip checks until it's about to decay (decay_due == 1) + if ( + pin.active + and pin.decay + and (pin.decay_due is not None) + and pin.decay_due > 1 + ): + continue + + # Include pin for checking when it has no decay, is inactive, or is about to decay + if (not pin.decay) or (not pin.active) or (pin.decay_due == 1): + pins_to_check[entry_id] = { + "condition": pin.condition, + "state": pin.condition_state, + } + + # Early return if nothing to check, but still tick decay + if not pins_to_check: + for entry_id, pin in world_state.pins.items(): + # Game-state-controlled pins do not decay + if pin.gamestate_condition: + continue + if pin.active and pin.decay: + if not pin.decay_due: + pin.decay_due = pin.decay + pin.decay_due -= self.check_pin_conditions_turns + log.debug("applying pin decay", pin=pin, decay_due=pin.decay_due) + if pin.decay_due <= 0: + log.debug("pin decay expired", pin=pin, decay_due=pin.decay_due) + pin.active = False + pin.decay_due = None + state_change = True + if state_change: + await self.scene.load_active_pins() + self.scene.emit_status() + return + + first_entry_id = list(pins_to_check.keys())[0] + + _, answers = await Prompt.request( + "world_state.check-pin-conditions", + self.client, + "analyze", + vars={ + "scene": self.scene, + "max_tokens": self.client.max_token_length, + "previous_states": json.dumps(pins_to_check, indent=2), + "coercion": {first_entry_id: {"condition": ""}}, + }, + ) + + # Apply LLM results + for entry_id, answer in answers.items(): + if entry_id not in world_state.pins: + log.debug( + "check_pin_conditions", + entry_id=entry_id, + answer=answer, + msg="entry_id not found in world_state.pins (LLM failed to produce a clean response)", + ) + continue + + log.debug("check_pin_conditions", entry_id=entry_id, answer=answer) + state = answer.get("state") + pin = world_state.pins[entry_id] + if state is True or ( + isinstance(state, str) and state.lower() in ["true", "yes", "y"] + ): + prev_state = pin.condition_state + pin.condition_state = True + if not pin.active: + state_change = True + pin.active = True + # Refresh decay countdown when condition is true and pin stays/turns active + if pin.decay: + pin.decay_due = pin.decay + if prev_state != pin.condition_state: + state_change = True + else: + if pin.condition_state is not False or pin.active: + pin.condition_state = False + pin.active = False + # Clear countdown when deactivated + pin.decay_due = None + state_change = True + + # Tick decay counters for all active pins with decay + for entry_id, pin in world_state.pins.items(): + # Game-state-controlled pins do not decay + if pin.gamestate_condition: + continue + if pin.active and pin.decay: + if not pin.decay_due: + pin.decay_due = pin.decay + # Decrement once per check cycle + pin.decay_due -= 1 + if pin.decay_due <= 0: + # Auto-deactivate on expiry + pin.active = False + pin.decay_due = None + state_change = True + + if state_change: + await self.scene.load_active_pins() + self.scene.emit_status() diff --git a/src/talemate/agents/world_state/reinforcements.py b/src/talemate/agents/world_state/reinforcements.py new file mode 100644 index 00000000..053545fa --- /dev/null +++ b/src/talemate/agents/world_state/reinforcements.py @@ -0,0 +1,197 @@ +from __future__ import annotations + +from typing import TYPE_CHECKING + +import structlog + +import talemate.emit.async_signals +from talemate.events import GameLoopEvent +from talemate.prompts import Prompt +from talemate.prompts.response import AnchorExtractor, ResponseSpec +from talemate.scene_message import ReinforcementMessage + +from talemate.agents.base import AgentAction, set_processing + +if TYPE_CHECKING: + from talemate.tale_mate import Character + +log = structlog.get_logger("talemate.agents.world_state") + + +class WorldStateReinforcementsMixin: + """ + World-state manager agent mixin that handles state reinforcements — the + periodic re-querying of tracked world/character questions and writing the + answers back into the scene (as detail, world entry, and/or inline message). + """ + + @classmethod + def add_actions(cls, actions: dict[str, AgentAction]): + actions["update_reinforcements"] = AgentAction( + enabled=True, + can_be_disabled=True, + container=True, + icon="mdi-image-auto-adjust", + label="Update state reinforcements", + description="Will attempt to update any due state reinforcements.", + config={}, + ) + + # config property helpers + + @property + def update_reinforcements_enabled(self) -> bool: + return self.resolve_enabled("update_reinforcements") + + # signal connect + + def connect(self, scene): + super().connect(scene) + talemate.emit.async_signals.get("game_loop").connect( + self.on_game_loop_update_reinforcements + ) + + async def on_game_loop_update_reinforcements(self, emission: GameLoopEvent): + """ + Called once per scene-loop round. + """ + if not self.enabled: + return + + await self.auto_update_reinforcments() + + # methods + + async def auto_update_reinforcments(self): + if not self.enabled: + return + + if not self.update_reinforcements_enabled: + return + + await self.update_reinforcements() + + @set_processing + async def update_reinforcements(self, force: bool = False, reset: bool = False): + """ + Queries due worldstate re-inforcements + """ + + for reinforcement in self.scene.world_state.reinforce: + # Skip character reinforcements if require_active is True and character is not active + if ( + reinforcement.require_active + and reinforcement.character + and not self.scene.character_is_active(reinforcement.character) + ): + continue + + if reinforcement.due <= 0 or force: + await self.update_reinforcement( + reinforcement.question, reinforcement.character, reset=reset + ) + else: + reinforcement.due -= 1 + + @set_processing + async def update_reinforcement( + self, question: str, character: "str | Character" = None, reset: bool = False + ) -> str: + """ + Queries a single re-inforcement + """ + + if isinstance(character, self.scene.Character): + character = character.name + + message = None + idx, reinforcement = await self.scene.world_state.find_reinforcement( + question, character + ) + + if not reinforcement: + log.warning( + "Reinforcement not found", question=question, character=character + ) + return + + message = ReinforcementMessage(message="") + message.set_source( + "world_state", + "update_reinforcement", + question=question, + character=character, + ) + + if reset and reinforcement.insert == "sequential": + self.scene.pop_history( + typ="reinforcement", meta_hash=message.meta_hash, all=True + ) + + if reinforcement.insert == "sequential": + kind = "analyze_freeform_medium_short" + else: + kind = "analyze_freeform" + + response, extracted = await Prompt.request( + "world_state.update-reinforcements", + self.client, + kind, + vars={ + "scene": self.scene, + "max_tokens": self.client.max_token_length, + "question": reinforcement.question, + "instructions": reinforcement.instructions or "", + "character": ( + self.scene.get_character(reinforcement.character) + if reinforcement.character + else None + ), + "answer": (reinforcement.answer if not reset else None) or "", + "reinforcement": reinforcement, + }, + response_spec=ResponseSpec( + extractors={ + "response": AnchorExtractor( + left="", + right="", + fallback_to_full=True, + ), + }, + ), + ) + + answer = extracted["response"] + + reinforcement.answer = answer + reinforcement.due = reinforcement.interval + + # remove any recent previous reinforcement message with same question + # to avoid overloading the near history with reinforcement messages + if not reset: + self.scene.pop_history( + typ="reinforcement", meta_hash=message.meta_hash, max_iterations=10 + ) + + if reinforcement.insert == "sequential": + # insert the reinforcement message at the current position + message.message = answer + log.debug("update_reinforcement", message=message, reset=reset) + await self.scene.push_history(message) + + # if reinforcement has a character name set, update the character detail + if reinforcement.character: + character = self.scene.get_character(reinforcement.character) + await character.set_detail(reinforcement.question, answer) + + else: + # set world entry + await self.scene.world_state_manager.save_world_entry( + reinforcement.question, + reinforcement.as_context_line, + {}, + ) + + self.scene.world_state.emit() + + return message