diff --git a/CLAUDE.md b/CLAUDE.md index 26b90a7..7138bd6 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -2,6 +2,7 @@ ## Commit Messages - NEVER add Claude attributions like "🤖 Generated with Claude Code" to commit messages +- NEVER add "Co-Authored-By: Claude " 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 diff --git a/activity_yaml_validator.py b/activity_yaml_validator.py index ea3b53c..c604c52 100644 --- a/activity_yaml_validator.py +++ b/activity_yaml_validator.py @@ -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): diff --git a/app.py b/app.py index 6c6386e..ea71569 100644 --- a/app.py +++ b/app.py @@ -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() diff --git a/research/activity29-battleship.yaml b/research/activity29-battleship.yaml index caf5753..cd42f99 100644 --- a/research/activity29-battleship.yaml +++ b/research/activity29-battleship.yaml @@ -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: diff --git a/research/activity29-testship.yaml b/research/activity29-testship.yaml index ec6d0a0..bf2556a 100644 --- a/research/activity29-testship.yaml +++ b/research/activity29-testship.yaml @@ -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: diff --git a/research/guarded_ai.py b/research/guarded_ai.py index f5e3723..1914b86 100644 --- a/research/guarded_ai.py +++ b/research/guarded_ai.py @@ -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", diff --git a/templates/chat.html b/templates/chat.html index ae62d5f..5ce6272 100644 --- a/templates/chat.html +++ b/templates/chat.html @@ -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 diff --git a/tests/unit/test_activity_yaml_validator.py b/tests/unit/test_activity_yaml_validator.py index 0397207..8c87fe9 100644 --- a/tests/unit/test_activity_yaml_validator.py +++ b/tests/unit/test_activity_yaml_validator.py @@ -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 diff --git a/tests/unit/test_app_feedback.py b/tests/unit/test_app_feedback.py new file mode 100644 index 0000000..61d80d9 --- /dev/null +++ b/tests/unit/test_app_feedback.py @@ -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) \ No newline at end of file diff --git a/tests/unit/test_guarded_ai.py b/tests/unit/test_guarded_ai.py new file mode 100644 index 0000000..0712e85 --- /dev/null +++ b/tests/unit/test_guarded_ai.py @@ -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) \ No newline at end of file