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