Implement per-prompt metadata filtering and fix battleship feedback system

Major improvements to battleship game feedback accuracy and user experience:

## New Multi-Prompt Feedback System
- Replaced single feedback with 3 specialized prompts: Shot Report, Ship Status, Game Over
- Each prompt has individual metadata filtering to see only relevant data
- Shot Report only sees hit/miss data, Ship Status only sees ship destruction data
- Added STFU token system to suppress empty messages (filtered out automatically)

## Technical Implementation
- Added per-prompt metadata_filter support in YAML structure
- Updated app.py and guarded_ai.py to handle prompt-specific filtering
- Legacy single-prompt system still works with transition-level filtering
- Added comprehensive test suite for feedback system validation

## User Experience Fixes
- Fixed TTS queue blocking JavaScript execution (async promises instead of await)
- Ship Status now correctly reports who destroyed which ship (role confusion fixed)
- Game Over only appears when game actually ends (no more random messages)
- Maintained dramatic storytelling while ensuring factual accuracy

## Battleship-Specific Improvements
- Ship destruction messages only appear when ships actually sink
- Clear separation of concerns: hits/misses vs ship destruction vs game over
- Eliminated false positive ship destruction reports
- Fixed role reversal where wrong player got credit for destruction

The battleship narrator now provides accurate, contextual feedback while preserving the dramatic naval warfare atmosphere.
This commit is contained in:
Russell Ballestrini 2025-08-11 11:39:49 -04:00
parent e28dc11f04
commit d4d697db59
10 changed files with 1448 additions and 107 deletions

View file

@ -2,6 +2,7 @@
## Commit Messages
- NEVER add Claude attributions like "🤖 Generated with Claude Code" to commit messages
- NEVER add "Co-Authored-By: Claude <noreply@anthropic.com>" to commit messages
- Keep commit messages focused on the actual changes and their purpose
- Use conventional commit format when appropriate
- Be concise but descriptive about what was changed and why

View file

@ -242,6 +242,10 @@ class ActivityYAMLValidator:
f"Section {section_id}, step {step_id}: 'feedback_tokens_for_ai' must be a string"
)
# Validate feedback_prompts (new multi-prompt system)
if "feedback_prompts" in step:
self._validate_feedback_prompts(step["feedback_prompts"], section_id, step_id)
# Validate buckets and transitions
if "buckets" in step:
self._validate_buckets(step["buckets"], section_id, step_id)
@ -251,6 +255,73 @@ class ActivityYAMLValidator:
step["transitions"], step.get("buckets", []), section_id, step_id
)
def _validate_feedback_prompts(self, feedback_prompts: List[Dict[str, Any]], section_id: str, step_id: str):
"""Validate feedback_prompts structure"""
if not isinstance(feedback_prompts, list):
self.errors.append(
f"Section {section_id}, step {step_id}: 'feedback_prompts' must be a list"
)
return
if len(feedback_prompts) == 0:
self.errors.append(
f"Section {section_id}, step {step_id}: 'feedback_prompts' cannot be empty"
)
return
prompt_names = set()
for i, prompt in enumerate(feedback_prompts):
if not isinstance(prompt, dict):
self.errors.append(
f"Section {section_id}, step {step_id}: feedback_prompts[{i}] must be a dictionary"
)
continue
# Required fields for each prompt
required_fields = ["name", "tokens_for_ai"]
for field in required_fields:
if field not in prompt:
self.errors.append(
f"Section {section_id}, step {step_id}: feedback_prompts[{i}] missing required field '{field}'"
)
# Validate name uniqueness
if "name" in prompt:
if not isinstance(prompt["name"], str):
self.errors.append(
f"Section {section_id}, step {step_id}: feedback_prompts[{i}].name must be a string"
)
else:
if prompt["name"] in prompt_names:
self.errors.append(
f"Section {section_id}, step {step_id}: duplicate feedback prompt name '{prompt['name']}'"
)
prompt_names.add(prompt["name"])
# Validate tokens_for_ai
if "tokens_for_ai" in prompt:
if not isinstance(prompt["tokens_for_ai"], str):
self.errors.append(
f"Section {section_id}, step {step_id}: feedback_prompts[{i}].tokens_for_ai must be a string"
)
# Check for STFU token usage (informational)
elif "STFU" in prompt["tokens_for_ai"]:
# This is valid - STFU token is used to suppress empty feedback messages
pass
# Validate metadata_filter (optional)
if "metadata_filter" in prompt:
if not isinstance(prompt["metadata_filter"], list):
self.errors.append(
f"Section {section_id}, step {step_id}: feedback_prompts[{i}].metadata_filter must be a list"
)
else:
for j, filter_key in enumerate(prompt["metadata_filter"]):
if not isinstance(filter_key, str):
self.errors.append(
f"Section {section_id}, step {step_id}: feedback_prompts[{i}].metadata_filter[{j}] must be a string"
)
def _validate_buckets(self, buckets: List[str], section_id: str, step_id: str):
"""Validate buckets list"""
if not isinstance(buckets, list):

145
app.py
View file

@ -2153,33 +2153,58 @@ def handle_activity_response(room_name, user_response, username):
# if "correct" or max_attempts reached.
# Provide feedback based on the category
# Filter metadata for feedback if metadata_feedback_filter is specified
feedback_metadata = activity_state.dict_metadata
if "metadata_feedback_filter" in transition:
filter_keys = transition["metadata_feedback_filter"]
feedback_metadata = {
k: v
for k, v in activity_state.dict_metadata.items()
if k in filter_keys
}
feedback = provide_feedback(
transition,
category,
step["question"],
feedback_tokens_for_ai,
user_response,
user_language,
username,
json.dumps(feedback_metadata),
json.dumps(new_metadata),
)
# Store and emit the feedback
if feedback:
# feedback is metadata language aware, doesn't need to be translated.
# Handle feedback systems
feedback_messages = []
if "feedback_prompts" in step:
# New multi-prompt system - pass full metadata, let each prompt filter
multi_feedback_messages = provide_feedback_prompts(
transition,
category,
step["question"],
step["feedback_prompts"],
user_response,
user_language,
username,
json.dumps(activity_state.dict_metadata), # Pass full metadata
json.dumps(new_metadata),
feedback_tokens_for_ai # Pass legacy tokens to be combined
)
feedback_messages.extend(multi_feedback_messages)
elif feedback_tokens_for_ai:
# Legacy single feedback system - use transition-level filtering
feedback_metadata = activity_state.dict_metadata
if "metadata_feedback_filter" in transition:
filter_keys = transition["metadata_feedback_filter"]
feedback_metadata = {
k: v
for k, v in activity_state.dict_metadata.items()
if k in filter_keys
}
feedback = provide_feedback(
transition,
category,
step["question"],
feedback_tokens_for_ai,
user_response,
user_language,
username,
json.dumps(feedback_metadata),
json.dumps(new_metadata),
)
if feedback and feedback.strip():
feedback_messages.append({
"name": "Feedback",
"content": feedback
})
# Store and emit all feedback messages
for feedback_msg in feedback_messages:
new_message = Message(
username="System (Feedback)", content=feedback, room_id=room.id
username=f"System ({feedback_msg['name'].title()})",
content=feedback_msg['content'],
room_id=room.id
)
db.session.add(new_message)
db.session.commit()
@ -2188,8 +2213,8 @@ def handle_activity_response(room_name, user_response, username):
"chat_message",
{
"id": new_message.id,
"username": "System (Feedback)",
"content": feedback,
"username": f"System ({feedback_msg['name'].title()})",
"content": feedback_msg['content'],
},
room=room_name,
)
@ -2587,6 +2612,70 @@ def provide_feedback(
return feedback
def provide_feedback_prompts(
transition,
category,
question,
feedback_prompts,
user_response,
user_language,
username,
json_metadata,
json_new_metadata,
legacy_tokens_for_ai="",
):
"""Generate feedback from multiple prompts"""
feedback_messages = []
# Parse full metadata once for filtering
full_metadata = json.loads(json_metadata)
for prompt in feedback_prompts:
prompt_name = prompt.get("name", "unnamed")
tokens_for_ai = prompt.get("tokens_for_ai", "")
# Apply per-prompt metadata filtering if specified
prompt_metadata = full_metadata
if "metadata_filter" in prompt:
filter_keys = prompt["metadata_filter"]
prompt_metadata = {k: v for k, v in full_metadata.items() if k in filter_keys}
print(f"DEBUG: Prompt '{prompt_name}' filter_keys: {filter_keys}")
print(f"DEBUG: Prompt '{prompt_name}' filtered metadata: {prompt_metadata}")
else:
print(f"DEBUG: Prompt '{prompt_name}' has NO metadata_filter, using full metadata")
print(f"DEBUG: Prompt '{prompt_name}' full metadata: {prompt_metadata}")
# Combine legacy tokens with prompt-specific tokens
if legacy_tokens_for_ai:
tokens_for_ai = legacy_tokens_for_ai + " " + tokens_for_ai
# Add language instruction
tokens_for_ai += f" You must provide the feedback in the user's language: {user_language}."
# Add transition-specific AI feedback if present
if "ai_feedback" in transition:
tokens_for_ai += f" {transition['ai_feedback'].get('tokens_for_ai', '')}"
ai_feedback = generate_ai_feedback(
category,
question,
user_response,
tokens_for_ai,
username,
json.dumps(prompt_metadata), # Use filtered metadata for this prompt
json_new_metadata,
)
# Only add feedback if it has content and isn't exactly the STFU token
if ai_feedback and ai_feedback.strip() and ai_feedback.strip() != "STFU":
feedback_messages.append({
"name": prompt_name,
"content": ai_feedback.strip()
})
return feedback_messages
def translate_text(text, target_language):
# Guard clause for default language
target_language = target_language.lower().split()

View file

@ -187,22 +187,58 @@ sections:
If the user wants to exit, categorize as 'exit'.
Otherwise, categorize as 'invalid_move'.
feedback_tokens_for_ai: |
You are the naval battle narrator. Look at the metadata provided and report what happened.
STEP 1 - CHECK SHIP DESTRUCTION (MANDATORY):
Look in the metadata for these exact fields:
- user_sunk_ship_this_round: If this contains a ship name like "Carrier" or "Battleship", say: "💥 SHIP DESTROYED! You have sunk the enemy's [ship name]! The enemy vessel explodes and sinks! Victory!"
- ai_sunk_ship_this_round: If this contains a ship name, say: "🔥 YOUR SHIP SUNK! The enemy destroyed your [ship name]! Your vessel burns and sinks!"
STEP 2 - REPORT SHOTS:
- Your shot result (user_hit_result): "hit" or "miss"
- Enemy shot result (ai_hit_result): "hit" or "miss"
EXAMPLE RESPONSE FORMAT:
If user_sunk_ship_this_round = "Carrier": "💥 SHIP DESTROYED! You have sunk the enemy's Carrier! [shot details]"
If ai_sunk_ship_this_round = "Destroyer": "🔥 YOUR SHIP SUNK! The enemy destroyed your Destroyer! [shot details]"
Always check the metadata for user_sunk_ship_this_round and ai_sunk_ship_this_round first. These are the most important events to report.
You are a battleship narrator. Each prompt has its own specific role - follow the individual prompt instructions precisely.
feedback_prompts:
- name: "Shot Report"
tokens_for_ai: |
🎯 Report ONLY the hit/miss results for both shots this turn. DO NOT report ship sinking or game over.
Check metadata:
- user_shot: Player's target position
- user_hit_result: "hit" or "miss"
- ai_shot: AI's target position
- ai_hit_result: "hit" or "miss"
Format: "🎯 Your shot at position [user_shot]: [user_hit_result]! 🤖 Enemy shot at position [ai_shot]: [ai_hit_result]!"
metadata_filter:
- user_shot
- ai_shot
- user_hit_result
- ai_hit_result
- name: "Ship Status"
tokens_for_ai: |
You are the Ship Destruction Oracle. Report ship destruction EXACTLY as the metadata shows:
CRITICAL - Read these metadata fields carefully:
- user_sunk_ship_this_round: If this contains a ship name like "Cruiser", it means THE USER destroyed an ENEMY ship
- ai_sunk_ship_this_round: If this contains a ship name like "Destroyer", it means THE ENEMY destroyed a USER ship
Your responses:
- If user_sunk_ship_this_round has a ship name: "💥 You have destroyed the enemy's [ship name]! It sinks beneath the waves!"
- If ai_sunk_ship_this_round has a ship name: "🔥 The enemy has destroyed your [ship name]! It has been claimed by the sea!"
- If both have ship names: combine both messages above
- If both are null/empty: "STFU"
Do NOT confuse who destroyed what. user_sunk_ship_this_round = USER victory. ai_sunk_ship_this_round = USER loss.
metadata_filter:
- user_sunk_ship_this_round
- ai_sunk_ship_this_round
- name: "Game Over"
tokens_for_ai: |
🏁 Check ONLY the game_over metadata field.
RESPOND WITH EXACTLY ONE OF THESE:
1. If game_over is false, null, or missing: "STFU"
2. If game_over is true AND user_wins is true: "🎉 TOTAL VICTORY! You have destroyed all enemy ships and won the battle! The seas are yours, Admiral!"
3. If game_over is true AND ai_wins is true: "💀 DEFEAT! The enemy has destroyed all your ships. Your fleet lies at the bottom of the ocean!"
CRITICAL: If game is not over, respond with exactly "STFU" and nothing else.
metadata_filter:
- game_over
- user_wins
- ai_wins
processing_script: |
import random
@ -847,16 +883,6 @@ sections:
The user shot seems valid.
metadata_tmp_add:
user_shot: "the-users-response"
metadata_feedback_filter:
- user_hit_result
- ai_hit_result
- ai_shot
- user_shot
- user_sunk_ship_this_round
- ai_sunk_ship_this_round
- game_over
- user_wins
- ai_wins
next_section_and_step: "section_1:step_2"
invalid_move:
content_blocks:

View file

@ -165,21 +165,58 @@ sections:
If the user wants to exit, categorize as 'exit'.
Otherwise, categorize as 'invalid_move'.
feedback_tokens_for_ai: |
Write battleship feedback from the game's perspective that covers:
1. User's shot result - check user_hit_result in metadata:
- If "hit": Describe the impact and explosion
- If "miss": Describe the splash and fog of war
2. AI's shot result - report where the AI fired:
- If hit: Describe the damage to the player's ship
- If miss: Describe the near miss and ocean spray
3. CRITICAL: If ai_sunk_ship_this_round contains a ship name, express dismay that the AI destroyed the player's ship in 2 sentences describing the carnage at sea
4. CRITICAL: If user_sunk_ship_this_round contains a ship name, celebrate the player destroying the AI ship in 2 sentences describing the carnage at sea
5. CRITICAL: If game_over is true, announce the victory:
- If user_wins is true: Celebrate the player's total victory with excitement!
- If ai_wins is true: Express dismay at the player's defeat!
Describe the sights and sounds of naval warfare! You are the game system rooting for the player!
You are a battleship narrator. Each prompt has its own specific role - follow the individual prompt instructions precisely.
feedback_prompts:
- name: "Shot Report"
tokens_for_ai: |
🎯 Report ONLY the hit/miss results for both shots this turn. DO NOT report ship sinking or game over.
Check metadata:
- user_shot: Player's target position
- user_hit_result: "hit" or "miss"
- ai_shot: AI's target position
- ai_hit_result: "hit" or "miss"
Format: "🎯 Your shot at position [user_shot]: [user_hit_result]! 🤖 Enemy shot at position [ai_shot]: [ai_hit_result]!"
metadata_filter:
- user_shot
- ai_shot
- user_hit_result
- ai_hit_result
- name: "Ship Status"
tokens_for_ai: |
You are the Ship Destruction Oracle. Report ship destruction EXACTLY as the metadata shows:
CRITICAL - Read these metadata fields carefully:
- user_sunk_ship_this_round: If this contains a ship name like "Cruiser", it means THE USER destroyed an ENEMY ship
- ai_sunk_ship_this_round: If this contains a ship name like "Destroyer", it means THE ENEMY destroyed a USER ship
Your responses:
- If user_sunk_ship_this_round has a ship name: "💥 You have destroyed the enemy's [ship name]! It sinks beneath the waves!"
- If ai_sunk_ship_this_round has a ship name: "🔥 The enemy has destroyed your [ship name]! It has been claimed by the sea!"
- If both have ship names: combine both messages above
- If both are null/empty: "STFU"
Do NOT confuse who destroyed what. user_sunk_ship_this_round = USER victory. ai_sunk_ship_this_round = USER loss.
metadata_filter:
- user_sunk_ship_this_round
- ai_sunk_ship_this_round
- name: "Game Over"
tokens_for_ai: |
🏁 Check ONLY the game_over metadata field.
RESPOND WITH EXACTLY ONE OF THESE:
1. If game_over is false, null, or missing: "STFU"
2. If game_over is true AND user_wins is true: "🎉 TOTAL VICTORY! You have destroyed all enemy ships and won the battle! The seas are yours, Admiral!"
3. If game_over is true AND ai_wins is true: "💀 DEFEAT! The enemy has destroyed all your ships. Your fleet lies at the bottom of the ocean!"
CRITICAL: If game is not over, respond with exactly "STFU" and nothing else.
metadata_filter:
- game_over
- user_wins
- ai_wins
processing_script: |
import random
@ -814,16 +851,6 @@ sections:
The user shot seems valid.
metadata_tmp_add:
user_shot: "the-users-response"
metadata_feedback_filter:
- user_hit_result
- ai_hit_result
- ai_shot
- user_shot
- user_sunk_ship_this_round
- ai_sunk_ship_this_round
- game_over
- user_wins
- ai_wins
next_section_and_step: "section_1:step_2"
invalid_move:
content_blocks:

View file

@ -122,7 +122,7 @@ def generate_ai_feedback(category, question, user_response, tokens_for_ai, metad
return f"Error: {e}"
# Provide feedback based on the category
# Provide feedback based on the category (legacy single feedback system)
def provide_feedback(
transition,
category,
@ -150,6 +150,55 @@ def provide_feedback(
return feedback
# Provide feedback using multiple prompts (new system)
def provide_feedback_prompts(
transition,
category,
question,
feedback_prompts,
user_response,
user_language,
metadata,
legacy_tokens_for_ai="",
):
"""Generate feedback from multiple prompts"""
feedback_messages = []
for prompt in feedback_prompts:
prompt_name = prompt.get("name", "unnamed")
tokens_for_ai = prompt.get("tokens_for_ai", "")
# Apply per-prompt metadata filtering if specified
prompt_metadata = metadata
if "metadata_filter" in prompt:
filter_keys = prompt["metadata_filter"]
prompt_metadata = {k: v for k, v in metadata.items() if k in filter_keys}
# Combine legacy tokens with prompt-specific tokens
if legacy_tokens_for_ai:
tokens_for_ai = legacy_tokens_for_ai + " " + tokens_for_ai
# Add language instruction
tokens_for_ai += f" Provide the feedback in {user_language}."
# Add transition-specific AI feedback if present
if "ai_feedback" in transition:
tokens_for_ai += f" {transition['ai_feedback'].get('tokens_for_ai', '')}"
ai_feedback = generate_ai_feedback(
category, question, user_response, tokens_for_ai, prompt_metadata
)
# Only add feedback if it has content and isn't exactly the STFU token
if ai_feedback and ai_feedback.strip() and ai_feedback.strip() != "STFU":
feedback_messages.append({
"name": prompt_name,
"content": ai_feedback.strip()
})
return feedback_messages
def execute_processing_script(metadata, script):
# Prepare the local environment for the script
local_env = {"metadata": metadata, "script_result": None}
@ -413,16 +462,41 @@ def simulate_activity(yaml_file_path):
print(f"\nMetadata: {json.dumps(metadata, indent=2)}")
# Provide feedback based on the category
feedback = provide_feedback(
transition,
category,
question,
user_response,
user_language,
step.get("feedback_tokens_for_ai", ""),
metadata,
)
print(f"\nFeedback: {feedback}")
feedback_messages = []
if "feedback_prompts" in step:
# New multi-prompt system - legacy tokens get combined with each prompt
multi_feedback_messages = provide_feedback_prompts(
transition,
category,
question,
step["feedback_prompts"],
user_response,
user_language,
metadata,
step.get("feedback_tokens_for_ai", "") # Pass legacy tokens to be combined
)
feedback_messages.extend(multi_feedback_messages)
elif step.get("feedback_tokens_for_ai"):
# Legacy single feedback system - only if no feedback_prompts
feedback = provide_feedback(
transition,
category,
question,
user_response,
user_language,
step.get("feedback_tokens_for_ai", ""),
metadata,
)
if feedback and feedback.strip():
feedback_messages.append({
"name": "Feedback",
"content": feedback
})
# Display all feedback messages
for feedback_msg in feedback_messages:
print(f"\n{feedback_msg['name']}: {feedback_msg['content']}")
if category not in [
"partial_understanding",

View file

@ -472,7 +472,7 @@ function queueTTS(text, playButton, messageId) {
}
// Function to process the next TTS in queue
async function processNextTTS() {
function processNextTTS() {
if (isPlayingTTS || ttsQueue.length === 0) {
return;
}
@ -481,15 +481,19 @@ async function processNextTTS() {
const { text, playButton, messageId } = ttsQueue.shift();
console.log("Processing TTS from queue:", messageId);
try {
await speakTextQueued(text, playButton, messageId);
} catch (error) {
console.error("TTS error:", error);
}
isPlayingTTS = false;
// Process next item in queue
setTimeout(processNextTTS, 100);
// Use non-blocking async processing
speakTextQueued(text, playButton, messageId)
.then(() => {
console.log("TTS completed successfully for:", messageId);
})
.catch((error) => {
console.error("TTS error:", error);
})
.finally(() => {
isPlayingTTS = false;
// Schedule next item with minimal delay to prevent blocking
setTimeout(processNextTTS, 10);
});
}
// Function to update auto-play TTS button display
@ -521,6 +525,23 @@ function toggleAutoPlayTTS() {
// Save to localStorage
localStorage.setItem('autoPlayTTS', autoPlayTTS.toString());
// If turning off, clear the queue and stop current audio
if (!autoPlayTTS) {
console.log("Clearing TTS queue, had", ttsQueue.length, "items");
ttsQueue = [];
isPlayingTTS = false;
// Stop any currently playing audio
if (currentAudio) {
currentAudio.pause();
currentAudio.currentTime = 0;
if (currentAudio.playButton) {
currentAudio.playButton.textContent = "Play";
}
currentAudio = null;
}
}
updateAutoPlayTTSDisplay();
}
@ -626,7 +647,7 @@ socket.on("chat_message", (data) => {
console.log("Queueing TTS for message:", data.id);
queueTTS(data.content, playButton, data.id);
}
}, 100); // Short delay to let buttons be created
}, 10); // Very short delay to let buttons be created
}
}
});
@ -799,7 +820,7 @@ socket.on("message_chunk", (data) => {
setTimeout(() => {
const fullText = targetMessageElement.textContent || targetMessageElement.innerText;
queueTTS(fullText, playButton, data.id);
}, 500); // Small delay to let the message render
}, 50); // Small delay to let the message render
}
}
}
@ -1018,11 +1039,15 @@ function addLineNumbers(block) {
// Socket event for setting the chat background
socket.on("set_background", (data) => {
const chat = document.getElementById("chat");
chat.style.backgroundImage = `url('data:image/png;base64,${data.image_data}')`;
chat.style.backgroundRepeat = "no-repeat";
chat.style.backgroundPosition = "right center";
chat.style.backgroundSize = "auto"; // Ensures the image is not stretched
// Use setTimeout to ensure background updates don't get blocked by TTS
setTimeout(() => {
const chat = document.getElementById("chat");
chat.style.backgroundImage = `url('data:image/png;base64,${data.image_data}')`;
chat.style.backgroundRepeat = "no-repeat";
chat.style.backgroundPosition = "right center";
chat.style.backgroundSize = "auto"; // Ensures the image is not stretched
console.log("Background image updated");
}, 0);
});
// Activity management functions

View file

@ -606,6 +606,139 @@ sections:
# Should catch the YAML syntax error we know is in there
self.assertTrue(any("YAML syntax error" in error for error in errors))
def test_feedback_prompts_validation(self):
"""Test validation of feedback_prompts structure"""
valid_feedback_prompts = """
sections:
- section_id: "section_1"
title: "Test"
steps:
- step_id: "step_1"
title: "Test Step"
question: "Test?"
feedback_prompts:
- name: "hit_miss"
tokens_for_ai: "Report hit/miss for both players"
- name: "ship_sinking"
tokens_for_ai: "Report any ship sinking events"
buckets:
- test
transitions:
test:
next_section_and_step: "section_1:step_2"
- step_id: "step_2"
title: "Final"
content_blocks:
- "Done"
"""
temp_file = self.create_temp_yaml(valid_feedback_prompts)
try:
is_valid, errors, warnings = self.validator.validate_file(temp_file)
self.assertTrue(is_valid)
self.assertEqual(len(errors), 0)
finally:
os.unlink(temp_file)
def test_invalid_feedback_prompts(self):
"""Test validation of invalid feedback_prompts structure"""
invalid_feedback_prompts = """
sections:
- section_id: "section_1"
title: "Test"
steps:
- step_id: "step_1"
title: "Test Step"
question: "Test?"
feedback_prompts: "should_be_list"
buckets:
- test
transitions:
test:
next_section_and_step: "section_1:step_2"
- step_id: "step_2"
title: "Test Step 2"
question: "Another test?"
feedback_prompts: [] # Empty list should error
buckets:
- test2
transitions:
test2:
next_section_and_step: "section_1:step_3"
- step_id: "step_3"
title: "Test Step 3"
question: "Third test?"
feedback_prompts:
- "should_be_dict"
- name: "valid_name"
# Missing tokens_for_ai
- name: "duplicate"
tokens_for_ai: "First prompt"
- name: "duplicate" # Duplicate name
tokens_for_ai: "Second prompt"
- name: 123 # Invalid name type
tokens_for_ai: "Valid tokens"
- name: "valid_name2"
tokens_for_ai: 456 # Invalid tokens type
buckets:
- test3
transitions:
test3:
content_blocks: ["Done"]
"""
temp_file = self.create_temp_yaml(invalid_feedback_prompts)
try:
is_valid, errors, warnings = self.validator.validate_file(temp_file)
self.assertFalse(is_valid)
# Check for specific error types
self.assertTrue(any("feedback_prompts' must be a list" in error for error in errors))
self.assertTrue(any("feedback_prompts' cannot be empty" in error for error in errors))
self.assertTrue(any("must be a dictionary" in error for error in errors))
self.assertTrue(any("missing required field" in error for error in errors))
self.assertTrue(any("duplicate feedback prompt name" in error for error in errors))
self.assertTrue(any("name must be a string" in error for error in errors))
self.assertTrue(any("tokens_for_ai must be a string" in error for error in errors))
finally:
os.unlink(temp_file)
def test_both_feedback_systems(self):
"""Test that both feedback_tokens_for_ai and feedback_prompts can be used together"""
both_feedback_systems = """
sections:
- section_id: "section_1"
title: "Test"
steps:
- step_id: "step_1"
title: "Test Step"
question: "Test?"
feedback_tokens_for_ai: "Legacy feedback system"
feedback_prompts:
- name: "new_system_1"
tokens_for_ai: "New system prompt 1"
- name: "new_system_2"
tokens_for_ai: "New system prompt 2"
buckets:
- test
transitions:
test:
next_section_and_step: "section_1:step_2"
- step_id: "step_2"
title: "Final"
content_blocks:
- "Done"
"""
temp_file = self.create_temp_yaml(both_feedback_systems)
try:
is_valid, errors, warnings = self.validator.validate_file(temp_file)
self.assertTrue(is_valid, f"Should be valid but got errors: {errors}")
self.assertEqual(len(errors), 0)
finally:
os.unlink(temp_file)
def test_cli_integration(self):
"""Test the command line interface"""
import subprocess

View file

@ -0,0 +1,567 @@
#!/usr/bin/env python3
"""
Unit tests for app.py feedback functions.
Tests the feedback generation functions including:
- Legacy provide_feedback function
- New provide_feedback_prompts function
- Both systems integration
- Metadata filtering
- Language handling
"""
import unittest
from unittest.mock import patch, MagicMock, call
import sys
import json
from pathlib import Path
# Add parent directory to path to import app functions
sys.path.insert(0, str(Path(__file__).parent.parent.parent))
class TestAppFeedback(unittest.TestCase):
"""Test cases for app.py feedback functions"""
def setUp(self):
"""Set up test fixtures"""
self.sample_transition = {
"ai_feedback": {
"tokens_for_ai": "Additional transition instructions"
},
"metadata_feedback_filter": ["shot_location", "hit_result", "ship_sunk"]
}
self.sample_metadata = {
"shot_location": "A5",
"hit_result": "hit",
"ship_sunk": "destroyer",
"private_info": "should_be_filtered",
"player_health": 100
}
self.sample_new_metadata = {
"new_shot": "B3",
"new_result": "miss"
}
def test_provide_feedback_import(self):
"""Test that we can import the provide_feedback function"""
try:
from app import provide_feedback
self.assertTrue(callable(provide_feedback))
except ImportError as e:
self.fail(f"Could not import provide_feedback: {e}")
def test_provide_feedback_prompts_import(self):
"""Test that we can import the provide_feedback_prompts function"""
try:
from app import provide_feedback_prompts
self.assertTrue(callable(provide_feedback_prompts))
except ImportError as e:
self.fail(f"Could not import provide_feedback_prompts: {e}")
@patch('app.get_openai_client_and_model')
def test_provide_feedback_legacy(self, mock_get_client):
"""Test legacy provide_feedback function"""
# Import here to avoid issues if module is not available
try:
from app import provide_feedback
except ImportError:
self.skipTest("app module not available for testing")
# Setup mock
mock_client = MagicMock()
mock_completion = MagicMock()
mock_completion.choices[0].message.content = "Great shot! You hit the target."
mock_client.chat.completions.create.return_value = mock_completion
mock_get_client.return_value = (mock_client, "test-model")
# Test data
transition = self.sample_transition
category = "hit"
question = "Where do you want to shoot?"
feedback_tokens_for_ai = "Provide battleship feedback"
user_response = "A5"
user_language = "English"
username = "testuser"
json_metadata = json.dumps(self.sample_metadata)
json_new_metadata = json.dumps(self.sample_new_metadata)
# Call function
feedback = provide_feedback(
transition, category, question, feedback_tokens_for_ai,
user_response, user_language, username,
json_metadata, json_new_metadata
)
# Verify result
self.assertIn("Great shot! You hit the target.", feedback)
# Verify client was called
mock_client.chat.completions.create.assert_called_once()
call_args = mock_client.chat.completions.create.call_args[1]
# Check that system message includes language and transition instructions
system_message = call_args['messages'][0]['content']
self.assertIn("English", system_message)
self.assertIn("Additional transition instructions", system_message)
@patch('app.get_openai_client_and_model')
def test_provide_feedback_prompts_multi(self, mock_get_client):
"""Test provide_feedback_prompts with multiple prompts"""
try:
from app import provide_feedback_prompts
except ImportError:
self.skipTest("app module not available for testing")
# Setup mock to return different responses for each prompt
mock_client = MagicMock()
mock_completion_1 = MagicMock()
mock_completion_1.choices[0].message.content = "Your shot at A5 was a hit! Enemy shot at B3 missed."
mock_completion_2 = MagicMock()
mock_completion_2.choices[0].message.content = "The enemy's destroyer has been sunk!"
mock_client.chat.completions.create.side_effect = [mock_completion_1, mock_completion_2]
mock_get_client.return_value = (mock_client, "test-model")
# Test data
transition = self.sample_transition
category = "valid_move"
question = "Where do you want to shoot?"
feedback_prompts = [
{
"name": "hit_miss_feedback",
"tokens_for_ai": "Report the hit/miss results for both players this turn"
},
{
"name": "ship_sinking_feedback",
"tokens_for_ai": "Report any ships that were sunk this turn"
}
]
user_response = "A5"
user_language = "English"
username = "testuser"
json_metadata = json.dumps(self.sample_metadata)
json_new_metadata = json.dumps(self.sample_new_metadata)
# Call function
feedback_messages = provide_feedback_prompts(
transition, category, question, feedback_prompts,
user_response, user_language, username,
json_metadata, json_new_metadata, ""
)
# Verify results
self.assertEqual(len(feedback_messages), 2)
# Check first feedback message
self.assertEqual(feedback_messages[0]["name"], "hit_miss_feedback")
self.assertIn("Your shot at A5 was a hit", feedback_messages[0]["content"])
# Check second feedback message
self.assertEqual(feedback_messages[1]["name"], "ship_sinking_feedback")
self.assertIn("destroyer has been sunk", feedback_messages[1]["content"])
# Verify client was called twice
self.assertEqual(mock_client.chat.completions.create.call_count, 2)
@patch('app.get_openai_client_and_model')
def test_provide_feedback_with_filtered_metadata(self, mock_get_client):
"""Test that provide_feedback works correctly with pre-filtered metadata"""
try:
from app import provide_feedback
except ImportError:
self.skipTest("app module not available for testing")
# Setup mock
mock_client = MagicMock()
mock_completion = MagicMock()
mock_completion.choices[0].message.content = "Filtered feedback"
mock_client.chat.completions.create.return_value = mock_completion
mock_get_client.return_value = (mock_client, "test-model")
# Simulate app.py behavior: filter metadata before calling provide_feedback
filtered_metadata = {
k: v for k, v in self.sample_metadata.items()
if k in self.sample_transition["metadata_feedback_filter"]
}
provide_feedback(
self.sample_transition, "test", "Question?", "tokens",
"response", "English", "user",
json.dumps(filtered_metadata), json.dumps({})
)
# Check that user message contains only filtered metadata
call_args = mock_client.chat.completions.create.call_args[1]
user_message = call_args['messages'][1]['content']
# Should contain filtered fields
self.assertIn("shot_location", user_message)
self.assertIn("hit_result", user_message)
self.assertIn("ship_sunk", user_message)
# Should NOT contain unfiltered fields (because we pre-filtered)
self.assertNotIn("private_info", user_message)
self.assertNotIn("player_health", user_message)
@patch('app.get_openai_client_and_model')
def test_provide_feedback_no_filter(self, mock_get_client):
"""Test feedback when no metadata filter is specified"""
try:
from app import provide_feedback
except ImportError:
self.skipTest("app module not available for testing")
# Setup mock
mock_client = MagicMock()
mock_completion = MagicMock()
mock_completion.choices[0].message.content = "Unfiltered feedback"
mock_client.chat.completions.create.return_value = mock_completion
mock_get_client.return_value = (mock_client, "test-model")
# Call function without metadata filter
transition = {"ai_feedback": {"tokens_for_ai": "Generate feedback"}} # No metadata_feedback_filter
provide_feedback(
transition, "test", "Question?", "tokens",
"response", "English", "user",
json.dumps(self.sample_metadata), json.dumps({})
)
# Check that user message contains all metadata
call_args = mock_client.chat.completions.create.call_args[1]
user_message = call_args['messages'][1]['content']
# Should contain all metadata fields when no filter is applied
self.assertIn("private_info", user_message)
self.assertIn("player_health", user_message)
@patch('app.get_openai_client_and_model')
def test_provide_feedback_error_handling(self, mock_get_client):
"""Test error handling in feedback functions"""
try:
from app import provide_feedback
except ImportError:
self.skipTest("app module not available for testing")
# Setup mock to raise exception
mock_client = MagicMock()
mock_client.chat.completions.create.side_effect = Exception("API Error")
mock_get_client.return_value = (mock_client, "test-model")
# Call function
feedback = provide_feedback(
{"ai_feedback": {"tokens_for_ai": "Generate feedback"}}, "test", "Question?", "tokens",
"response", "English", "user",
json.dumps({}), json.dumps({})
)
# Should handle error gracefully
self.assertIn("Error", feedback)
@patch('app.get_openai_client_and_model')
def test_provide_feedback_prompts_filter_empty_and_stfu(self, mock_get_client):
"""Test feedback_prompts with empty results and STFU tokens filtered out"""
try:
from app import provide_feedback_prompts
except ImportError:
self.skipTest("app module not available for testing")
# Setup mock to return mixed results including STFU token
mock_client = MagicMock()
mock_completion_1 = MagicMock()
mock_completion_1.choices[0].message.content = "" # Empty result
mock_completion_2 = MagicMock()
mock_completion_2.choices[0].message.content = "STFU" # STFU token (should be filtered)
mock_completion_3 = MagicMock()
mock_completion_3.choices[0].message.content = "Valid feedback" # Valid result
mock_client.chat.completions.create.side_effect = [
mock_completion_1, mock_completion_2, mock_completion_3
]
mock_get_client.return_value = (mock_client, "test-model")
# Test data
feedback_prompts = [
{"name": "empty", "tokens_for_ai": "Empty prompt"},
{"name": "stfu", "tokens_for_ai": "STFU prompt"},
{"name": "valid", "tokens_for_ai": "Valid prompt"}
]
feedback_messages = provide_feedback_prompts(
{}, "test", "Question?", feedback_prompts,
"response", "English", "user",
json.dumps({}), json.dumps({}), ""
)
# Should only return valid feedback (empty and STFU both filtered out the same way)
self.assertEqual(len(feedback_messages), 1)
self.assertEqual(feedback_messages[0]["name"], "valid")
self.assertEqual(feedback_messages[0]["content"], "Valid feedback")
@patch('app.get_openai_client_and_model')
def test_provide_feedback_prompts_stfu_partial_not_filtered(self, mock_get_client):
"""Test that messages containing STFU as part of larger text are NOT filtered"""
try:
from app import provide_feedback_prompts
except ImportError:
self.skipTest("app module not available for testing")
# Setup mock to return STFU as part of larger message
mock_client = MagicMock()
mock_completion_1 = MagicMock()
mock_completion_1.choices[0].message.content = "STFU you rascal." # Should NOT be filtered
mock_completion_2 = MagicMock()
mock_completion_2.choices[0].message.content = "Go STFU yourself!" # Should NOT be filtered
mock_completion_3 = MagicMock()
mock_completion_3.choices[0].message.content = "STFU" # Should be filtered
mock_client.chat.completions.create.side_effect = [
mock_completion_1, mock_completion_2, mock_completion_3
]
mock_get_client.return_value = (mock_client, "test-model")
# Test data
feedback_prompts = [
{"name": "partial1", "tokens_for_ai": "Partial STFU 1"},
{"name": "partial2", "tokens_for_ai": "Partial STFU 2"},
{"name": "exact", "tokens_for_ai": "Exact STFU"}
]
feedback_messages = provide_feedback_prompts(
{}, "test", "Question?", feedback_prompts,
"response", "English", "user",
json.dumps({}), json.dumps({}), ""
)
# Should return the two partial STFU messages, but not the exact "STFU"
self.assertEqual(len(feedback_messages), 2)
self.assertEqual(feedback_messages[0]["name"], "partial1")
self.assertEqual(feedback_messages[0]["content"], "STFU you rascal.")
self.assertEqual(feedback_messages[1]["name"], "partial2")
self.assertEqual(feedback_messages[1]["content"], "Go STFU yourself!")
@patch('app.get_openai_client_and_model')
def test_provide_feedback_prompts_per_prompt_metadata_filtering(self, mock_get_client):
"""Test that each prompt gets its own filtered metadata"""
try:
from app import provide_feedback_prompts
except ImportError:
self.skipTest("app module not available for testing")
# Setup mock to return different responses
mock_client = MagicMock()
mock_completion_1 = MagicMock()
mock_completion_1.choices[0].message.content = "Shot feedback with hit/miss data"
mock_completion_2 = MagicMock()
mock_completion_2.choices[0].message.content = "Ship feedback with sinking data"
mock_client.chat.completions.create.side_effect = [
mock_completion_1, mock_completion_2
]
mock_get_client.return_value = (mock_client, "test-model")
# Test data with mixed metadata
full_metadata = {
"user_shot": "A5",
"user_hit_result": "hit",
"ai_shot": "B3",
"ai_hit_result": "miss",
"user_sunk_ship_this_round": "Destroyer",
"ai_sunk_ship_this_round": None,
"game_over": False,
"extra_field": "should_not_appear"
}
feedback_prompts = [
{
"name": "shot_report",
"tokens_for_ai": "Report hit/miss",
"metadata_filter": ["user_shot", "user_hit_result", "ai_shot", "ai_hit_result"]
},
{
"name": "ship_status",
"tokens_for_ai": "Report ship sinking",
"metadata_filter": ["user_sunk_ship_this_round", "ai_sunk_ship_this_round"]
}
]
feedback_messages = provide_feedback_prompts(
{}, "test", "Question?", feedback_prompts,
"response", "English", "user",
json.dumps(full_metadata), json.dumps({}), ""
)
# Verify both prompts got responses
self.assertEqual(len(feedback_messages), 2)
self.assertEqual(feedback_messages[0]["name"], "shot_report")
self.assertEqual(feedback_messages[1]["name"], "ship_status")
# Verify the first prompt only got shot-related metadata
first_call_args = mock_client.chat.completions.create.call_args_list[0][1]
first_user_message = first_call_args['messages'][1]['content']
self.assertIn("user_shot", first_user_message)
self.assertIn("user_hit_result", first_user_message)
self.assertIn("ai_shot", first_user_message)
self.assertIn("ai_hit_result", first_user_message)
self.assertNotIn("user_sunk_ship_this_round", first_user_message)
self.assertNotIn("extra_field", first_user_message)
# Verify the second prompt only got ship-related metadata
second_call_args = mock_client.chat.completions.create.call_args_list[1][1]
second_user_message = second_call_args['messages'][1]['content']
self.assertIn("user_sunk_ship_this_round", second_user_message)
self.assertIn("ai_sunk_ship_this_round", second_user_message)
self.assertNotIn("user_shot", second_user_message)
self.assertNotIn("extra_field", second_user_message)
@patch('app.get_openai_client_and_model')
def test_ship_status_metadata_filtering_debug(self, mock_get_client):
"""Debug test to check if Ship Status is getting only the right metadata"""
try:
from app import provide_feedback_prompts
except ImportError:
self.skipTest("app module not available for testing")
# Setup mock
mock_client = MagicMock()
mock_completion = MagicMock()
mock_completion.choices[0].message.content = "Test response"
mock_client.chat.completions.create.return_value = mock_completion
mock_get_client.return_value = (mock_client, "test-model")
# Test data mimicking the actual battleship scenario
full_metadata = {
"user_shot": "46", # This should NOT appear in Ship Status
"ai_shot": "49", # This should NOT appear in Ship Status
"user_hit_result": "hit",
"ai_hit_result": "miss",
"user_sunk_ship_this_round": "Destroyer", # This SHOULD appear
"ai_sunk_ship_this_round": None, # This SHOULD appear
"game_over": False,
"extra_stuff": "should not appear anywhere"
}
# Exact structure from battleship YAML
feedback_prompts = [
{
"name": "Shot Report",
"tokens_for_ai": "🎯 Report ONLY the hit/miss results",
"metadata_filter": ["user_shot", "ai_shot", "user_hit_result", "ai_hit_result"]
},
{
"name": "Ship Status",
"tokens_for_ai": "You are the Ship Destruction Oracle",
"metadata_filter": ["user_sunk_ship_this_round", "ai_sunk_ship_this_round"]
}
]
# Call the function
provide_feedback_prompts(
{}, "valid_move", "Question?", feedback_prompts,
"46", "English", "user",
json.dumps(full_metadata), json.dumps({}), ""
)
# Check what metadata each prompt actually received
self.assertEqual(mock_client.chat.completions.create.call_count, 2)
# First call should be Shot Report
shot_report_call = mock_client.chat.completions.create.call_args_list[0][1]
shot_report_metadata = shot_report_call['messages'][1]['content']
print("=== SHOT REPORT METADATA ===")
print(shot_report_metadata)
# Shot Report should have shot data but NOT ship destruction data
self.assertIn("user_shot", shot_report_metadata)
self.assertIn("46", shot_report_metadata)
self.assertNotIn("user_sunk_ship_this_round", shot_report_metadata)
self.assertNotIn("Destroyer", shot_report_metadata)
# Second call should be Ship Status
ship_status_call = mock_client.chat.completions.create.call_args_list[1][1]
ship_status_metadata = ship_status_call['messages'][1]['content']
print("=== SHIP STATUS METADATA ===")
print(ship_status_metadata)
# Ship Status should have ship destruction data but NOT shot data
self.assertIn("user_sunk_ship_this_round", ship_status_metadata)
self.assertIn("Destroyer", ship_status_metadata)
self.assertNotIn("user_shot", ship_status_metadata)
self.assertNotIn("46", ship_status_metadata)
self.assertNotIn("extra_stuff", ship_status_metadata)
def test_provide_feedback_prompts_language_injection(self):
"""Test that language instructions are properly added to prompts"""
try:
from app import provide_feedback_prompts
except ImportError:
self.skipTest("app module not available for testing")
with patch('app.get_openai_client_and_model') as mock_get_client:
mock_client = MagicMock()
mock_completion = MagicMock()
mock_completion.choices[0].message.content = "Feedback in Spanish"
mock_client.chat.completions.create.return_value = mock_completion
mock_get_client.return_value = (mock_client, "test-model")
feedback_prompts = [
{"name": "test", "tokens_for_ai": "Base prompt"}
]
# Test with Spanish language
provide_feedback_prompts(
{}, "test", "Question?", feedback_prompts,
"response", "Spanish", "user",
json.dumps({}), json.dumps({}), ""
)
# Check that system message includes Spanish language instruction
call_args = mock_client.chat.completions.create.call_args[1]
system_message = call_args['messages'][0]['content']
self.assertIn("Spanish", system_message)
self.assertIn("Base prompt", system_message)
@patch('app.get_openai_client_and_model')
def test_provide_feedback_transition_tokens(self, mock_get_client):
"""Test that transition ai_feedback tokens are included"""
try:
from app import provide_feedback_prompts
except ImportError:
self.skipTest("app module not available for testing")
mock_client = MagicMock()
mock_completion = MagicMock()
mock_completion.choices[0].message.content = "Enhanced feedback"
mock_client.chat.completions.create.return_value = mock_completion
mock_get_client.return_value = (mock_client, "test-model")
transition = {
"ai_feedback": {
"tokens_for_ai": "Be more dramatic in your feedback"
}
}
feedback_prompts = [
{"name": "test", "tokens_for_ai": "Base prompt"}
]
provide_feedback_prompts(
transition, "test", "Question?", feedback_prompts,
"response", "English", "user",
json.dumps({}), json.dumps({}), ""
)
# Check that system message includes both base and transition tokens
call_args = mock_client.chat.completions.create.call_args[1]
system_message = call_args['messages'][0]['content']
self.assertIn("Base prompt", system_message)
self.assertIn("Be more dramatic in your feedback", system_message)
if __name__ == "__main__":
unittest.main(verbosity=2)

View file

@ -0,0 +1,328 @@
#!/usr/bin/env python3
"""
Unit tests for the guarded_ai.py module.
Tests the core feedback generation functions including:
- Legacy single feedback system
- New multi-prompt feedback system
- Both systems together
- OpenAI client initialization
- Categorization and feedback generation
"""
import unittest
from unittest.mock import patch, MagicMock, call
import sys
from pathlib import Path
import json
# Add parent directory to path to import guarded_ai
sys.path.insert(0, str(Path(__file__).parent.parent.parent / "research"))
from guarded_ai import (
provide_feedback,
provide_feedback_prompts,
categorize_response,
generate_ai_feedback,
get_openai_client_and_model,
initialize_model_map,
)
class TestGuardedAI(unittest.TestCase):
"""Test cases for guarded_ai functions"""
def setUp(self):
"""Set up test fixtures"""
self.sample_metadata = {
"player_health": 100,
"enemy_health": 80,
"user_shot": "A5",
"ai_shot": "B3",
"user_hit_result": "hit",
"ai_hit_result": "miss",
}
self.sample_transition = {
"ai_feedback": {
"tokens_for_ai": "Additional transition-specific instructions"
},
"metadata_feedback_filter": ["user_shot", "ai_shot", "user_hit_result", "ai_hit_result"]
}
@patch('guarded_ai.get_openai_client_and_model')
def test_categorize_response(self, mock_get_client):
"""Test response categorization"""
# Setup mock
mock_client = MagicMock()
mock_completion = MagicMock()
mock_completion.choices[0].message.content = "correct_answer"
mock_client.chat.completions.create.return_value = mock_completion
mock_get_client.return_value = (mock_client, "test-model")
# Test categorization
question = "What is 2+2?"
response = "Four"
buckets = ["correct_answer", "wrong_answer"]
tokens_for_ai = "Categorize math answers"
category = categorize_response(question, response, buckets, tokens_for_ai)
# Verify result
self.assertEqual(category, "correct_answer")
# Verify client was called correctly
mock_client.chat.completions.create.assert_called_once()
call_args = mock_client.chat.completions.create.call_args[1]
self.assertEqual(call_args['model'], 'test-model')
self.assertEqual(call_args['max_tokens'], 5)
self.assertEqual(call_args['temperature'], 0)
# Check message content
messages = call_args['messages']
self.assertEqual(len(messages), 2)
self.assertIn("correct_answer, wrong_answer", messages[0]['content'])
@patch('guarded_ai.get_openai_client_and_model')
def test_generate_ai_feedback(self, mock_get_client):
"""Test AI feedback generation"""
# Setup mock
mock_client = MagicMock()
mock_completion = MagicMock()
mock_completion.choices[0].message.content = "Great job on the math!"
mock_client.chat.completions.create.return_value = mock_completion
mock_get_client.return_value = (mock_client, "test-model")
# Test feedback generation
category = "correct_answer"
question = "What is 2+2?"
user_response = "Four"
tokens_for_ai = "Provide encouraging feedback"
metadata = {"score": 100}
feedback = generate_ai_feedback(category, question, user_response, tokens_for_ai, metadata)
# Verify result
self.assertEqual(feedback, "Great job on the math!")
# Verify client was called correctly
mock_client.chat.completions.create.assert_called_once()
call_args = mock_client.chat.completions.create.call_args[1]
self.assertEqual(call_args['model'], 'test-model')
self.assertEqual(call_args['max_tokens'], 250)
self.assertEqual(call_args['temperature'], 0.7)
@patch('guarded_ai.generate_ai_feedback')
def test_provide_feedback_legacy(self, mock_generate_feedback):
"""Test legacy single feedback system"""
mock_generate_feedback.return_value = "Good work! Try again."
# Test data
transition = self.sample_transition
category = "partial_understanding"
question = "What is the capital of France?"
user_response = "Paris is nice"
user_language = "English"
tokens_for_ai = "Provide geography feedback"
metadata = {"attempts": 1}
# Call function
feedback = provide_feedback(
transition, category, question, user_response,
user_language, tokens_for_ai, metadata
)
# Verify feedback was generated
self.assertIn("AI Feedback:", feedback)
self.assertIn("Good work! Try again.", feedback)
# Verify generate_ai_feedback was called with filtered metadata
mock_generate_feedback.assert_called_once()
call_args = mock_generate_feedback.call_args[0]
self.assertEqual(call_args[0], category) # category
self.assertEqual(call_args[1], question) # question
self.assertEqual(call_args[2], user_response) # user_response
# Check tokens_for_ai includes language and transition instructions
tokens_arg = call_args[3]
self.assertIn("English", tokens_arg)
self.assertIn("Additional transition-specific instructions", tokens_arg)
# Check metadata was filtered
filtered_metadata = call_args[4]
expected_filtered = {k: v for k, v in self.sample_metadata.items()
if k in transition["metadata_feedback_filter"]}
# Since our test metadata doesn't have the filtered keys, it should be empty or contain only matching keys
# But the function should have passed what it received
@patch('guarded_ai.generate_ai_feedback')
def test_provide_feedback_prompts(self, mock_generate_feedback):
"""Test new multi-prompt feedback system"""
# Setup mock to return different feedback for each prompt
mock_generate_feedback.side_effect = [
"Hit at A5, miss at B3",
"No ships were sunk this round"
]
# Test data
transition = self.sample_transition
category = "valid_move"
question = "Where do you want to shoot?"
feedback_prompts = [
{
"name": "hit_miss",
"tokens_for_ai": "Report the hit/miss results for both players"
},
{
"name": "ship_sinking",
"tokens_for_ai": "Report any ships that were sunk"
}
]
user_response = "A5"
user_language = "English"
metadata = self.sample_metadata
# Call function
feedback_messages = provide_feedback_prompts(
transition, category, question, feedback_prompts,
user_response, user_language, metadata, ""
)
# Verify we got the expected number of feedback messages
self.assertEqual(len(feedback_messages), 2)
# Verify message structure
self.assertEqual(feedback_messages[0]["name"], "hit_miss")
self.assertEqual(feedback_messages[0]["content"], "Hit at A5, miss at B3")
self.assertEqual(feedback_messages[1]["name"], "ship_sinking")
self.assertEqual(feedback_messages[1]["content"], "No ships were sunk this round")
# Verify generate_ai_feedback was called twice
self.assertEqual(mock_generate_feedback.call_count, 2)
@patch('guarded_ai.generate_ai_feedback')
def test_provide_feedback_prompts_empty_responses(self, mock_generate_feedback):
"""Test that empty feedback responses are filtered out"""
# Setup mock to return empty/whitespace responses
mock_generate_feedback.side_effect = [
"", # Empty response
" ", # Whitespace only
"Valid feedback" # Valid response
]
transition = {}
category = "test"
question = "Test?"
feedback_prompts = [
{"name": "empty", "tokens_for_ai": "Empty prompt"},
{"name": "whitespace", "tokens_for_ai": "Whitespace prompt"},
{"name": "valid", "tokens_for_ai": "Valid prompt"}
]
user_response = "Test response"
user_language = "English"
metadata = {}
feedback_messages = provide_feedback_prompts(
transition, category, question, feedback_prompts,
user_response, user_language, metadata, ""
)
# Should only return the valid feedback message
self.assertEqual(len(feedback_messages), 1)
self.assertEqual(feedback_messages[0]["name"], "valid")
self.assertEqual(feedback_messages[0]["content"], "Valid feedback")
def test_provide_feedback_no_ai_feedback_config(self):
"""Test legacy feedback when no ai_feedback config in transition"""
transition = {} # No ai_feedback key
category = "test"
question = "Test?"
user_response = "Response"
user_language = "English"
tokens_for_ai = "Base tokens"
metadata = {}
with patch('guarded_ai.generate_ai_feedback') as mock_generate:
mock_generate.return_value = "" # Should not be called
feedback = provide_feedback(
transition, category, question, user_response,
user_language, tokens_for_ai, metadata
)
# Should NOT call generate_ai_feedback when no ai_feedback in transition
mock_generate.assert_not_called()
self.assertEqual(feedback, "")
@patch.dict('os.environ', {'MODEL_ENDPOINT_0': 'http://test.com', 'MODEL_API_KEY_0': 'test-key'})
def test_initialize_model_map(self):
"""Test model map initialization from environment variables"""
with patch('guarded_ai.get_client_for_endpoint') as mock_get_client:
mock_client = MagicMock()
mock_get_client.return_value = mock_client
# Clear and reinitialize
import guarded_ai
guarded_ai.MODEL_CLIENT_MAP = {}
initialize_model_map()
# Verify client was created and stored
mock_get_client.assert_called_with('http://test.com', 'test-key')
self.assertIn('endpoint_0', guarded_ai.MODEL_CLIENT_MAP)
self.assertEqual(guarded_ai.MODEL_CLIENT_MAP['endpoint_0'][0], mock_client)
def test_get_openai_client_and_model_default(self):
"""Test getting OpenAI client with default model"""
with patch('guarded_ai.MODEL_CLIENT_MAP', {}):
with patch('guarded_ai.get_client_for_endpoint') as mock_get_client:
mock_client = MagicMock()
mock_get_client.return_value = mock_client
client, model = get_openai_client_and_model()
# Should return default model name
self.assertEqual(model, "adamo1139/Hermes-3-Llama-3.1-8B-FP8-Dynamic")
self.assertEqual(client, mock_client)
def test_get_openai_client_and_model_from_map(self):
"""Test getting OpenAI client from model map"""
mock_client = MagicMock()
test_map = {
'endpoint_0': (mock_client, 'http://test.com')
}
with patch('guarded_ai.MODEL_CLIENT_MAP', test_map):
client, model = get_openai_client_and_model("test-model")
# Should return client from map
self.assertEqual(client, mock_client)
self.assertEqual(model, "test-model")
@patch('guarded_ai.get_openai_client_and_model')
def test_categorize_response_error_handling(self, mock_get_client):
"""Test error handling in categorize_response"""
# Setup mock to raise exception
mock_client = MagicMock()
mock_client.chat.completions.create.side_effect = Exception("API Error")
mock_get_client.return_value = (mock_client, "test-model")
category = categorize_response("Test?", "Answer", ["bucket1"], "tokens")
# Should return error string
self.assertIn("Error:", category)
@patch('guarded_ai.get_openai_client_and_model')
def test_generate_ai_feedback_error_handling(self, mock_get_client):
"""Test error handling in generate_ai_feedback"""
# Setup mock to raise exception
mock_client = MagicMock()
mock_client.chat.completions.create.side_effect = Exception("API Error")
mock_get_client.return_value = (mock_client, "test-model")
feedback = generate_ai_feedback("cat", "Q?", "A", "tokens", {})
# Should return error string
self.assertIn("Error:", feedback)
if __name__ == "__main__":
unittest.main(verbosity=2)