mirror of
https://github.com/vegu-ai/talemate.git
synced 2026-09-01 19:48:52 +02:00
Refactor world state management by introducing mixins for reinforcements and pin conditions
This commit is contained in:
@@ -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="<ANSWER>",
|
||||
right="</ANSWER>",
|
||||
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:
|
||||
"""
|
||||
|
||||
224
src/talemate/agents/world_state/pin_conditions.py
Normal file
224
src/talemate/agents/world_state/pin_conditions.py
Normal file
@@ -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()
|
||||
197
src/talemate/agents/world_state/reinforcements.py
Normal file
197
src/talemate/agents/world_state/reinforcements.py
Normal file
@@ -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="<ANSWER>",
|
||||
right="</ANSWER>",
|
||||
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
|
||||
Reference in New Issue
Block a user