diff --git a/README.md b/README.md index 9fe920c..1bf797c 100644 --- a/README.md +++ b/README.md @@ -31,14 +31,15 @@ This will start the interactive setup. Here's what you'll need: - Install [plugin](https://github.com/coddingtonbear/obsidian-local-rest-api) in Obsidian - Copy the API key from plugin settings -2. **Anthropic API key:** - - Get from [console.anthropic.com](https://console.anthropic.com/) +2. **AI Model:** + - Choose from popular models (Claude 4, GPT-4o, Gemini 2.0, etc.) + - 8+ providers supported through LiteLLM 3. **AnkiConnect setup:** - Add-on code: `2055492159` - Keep Anki running -You can then follow the interactive setup and edit the configuration as you like. +The interactive setup will guide you through model selection and configuration. ## Usage @@ -49,9 +50,10 @@ oki # Alias ### Configuration ```bash -oki config # Show config -oki config get max_cards # Get specific setting -oki config set max_cards 15 # Update setting +oki config # Show config +oki config get max_cards # Get specific setting +oki config set max_cards 15 # Update setting +oki config set model "GPT-5" # Switch AI model ``` ### Tags diff --git a/obsidianki/ai/client.py b/obsidianki/ai/client.py index 329a91b..3be203c 100644 --- a/obsidianki/ai/client.py +++ b/obsidianki/ai/client.py @@ -1,33 +1,50 @@ import os -from anthropic import Anthropic -from typing import List, Dict, cast, Any +from typing import List, Dict, Optional, Union, cast +import litellm +from litellm import completion +from litellm.types.utils import ModelResponse from obsidianki.cli.config import console, CONFIG from obsidianki.cli.utils import process_code_blocks, strip_html from obsidianki.cli.models import Note, Flashcard -from anthropic.types import ToolChoiceParam, MessageParam +from obsidianki.ai.models import MODEL_MAP from obsidianki.ai.prompts import SYSTEM_PROMPT, QUERY_SYSTEM_PROMPT, TARGETED_SYSTEM_PROMPT, MULTI_TURN_DQL_AGENT_PROMPT from obsidianki.ai.tools import FLASHCARD_TOOL, DQL_EXECUTION_TOOL, FINALIZE_SELECTION_TOOL AI_RESULT_SET_SIZE = 20 +# Suppress litellm logging +litellm.suppress_debug_info = True + class FlashcardAI: def __init__(self): - self.client = Anthropic(api_key=os.getenv("ANTHROPIC_API_KEY")) + model_name = getattr(CONFIG, 'model', 'Claude Sonnet 4.5') + + model_info = MODEL_MAP[model_name] + self.provider = model_info["provider"] + self.model = model_info["model"] + + self._validate_api_key() + + def _validate_api_key(self) -> None: + """Ensure appropriate API key is available for selected provider""" + key_map = {model_info["provider"]: model_info["key_name"] for model_info in MODEL_MAP.values()} + + required_key = key_map.get(self.provider) + + if required_key is None: + raise ValueError(f"{required_key} not found in environment variables for provider {self.provider}") - if not os.getenv("ANTHROPIC_API_KEY"): - raise ValueError("ANTHROPIC_API_KEY not found in environment variables") - def _build_card_instruction(self, target_cards: int) -> str: context = f"create approximately {target_cards} flashcards." if CONFIG.use_extrapolation: context += " IMPORTANT: You are allowed to extrapolate with your pre-existing knowledge somewhat if you feel it is directly relevant to note substance, but is not written in the note itself." return context - + def _build_dedup_context(self, previous_fronts: List[str]) -> str: if not previous_fronts: return "" - + previous_questions = "\n".join([f"- {front}" for front in previous_fronts]) dedup_context = f""" @@ -35,7 +52,7 @@ def _build_dedup_context(self, previous_fronts: List[str]) -> str: {previous_questions} DO NOT create flashcards that ask similar questions or cover the same concepts as the ones listed above. Focus on different aspects of the content.""" - + return dedup_context def _build_schema_context(self, deck_examples: List[Dict[str, str]]) -> str: @@ -116,9 +133,44 @@ def _build_difficulty_context(self) -> str: return "" - def generate_flashcards(self, note: Note, target_cards: int, previous_fronts: list = [], deck_examples: list = []) -> List[Flashcard]: - """Generate flashcards from a Note object using Claude""" - + def _get_tool_choice(self, function_name: str) -> str: + return 'required' + + def _call_llm( + self, + system_prompt: str, + user_prompt: str, + tools: List[Dict[str, object]], + tool_choice: Union[str, Dict[str, object]], + max_tokens: int = 8000 + ) -> Optional[ModelResponse]: + """Unified LLM call using litellm""" + try: + response = completion( + model=self.model, + messages=[ + {"role": "system", "content": system_prompt}, + {"role": "user", "content": user_prompt} + ], + tools=tools, + tool_choice=tool_choice, + max_tokens=max_tokens + ) + # We never use streaming, so response is always ModelResponse + return cast(ModelResponse, response) + except Exception as e: + import traceback + console.print(f"[red]ERROR:[/red] LLM call failed" + str(e)) + return None + + def generate_flashcards( + self, + note: Note, + target_cards: int, + previous_fronts: List[str] = [], + deck_examples: List[Dict[str, str]] = [] + ) -> List[Flashcard]: + """Generate flashcards from a Note object using LLM""" card_instruction = self._build_card_instruction(target_cards) dedup_context = self._build_dedup_context(previous_fronts) schema_context = self._build_schema_context(deck_examples) @@ -131,52 +183,59 @@ def generate_flashcards(self, note: Note, target_cards: int, previous_fronts: li Please analyze this note and {card_instruction} for the key information that would be valuable for spaced repetition learning.""" - try: - response = self.client.messages.create( - model="claude-4-sonnet-20250514", - max_tokens=8000, - system=SYSTEM_PROMPT, - messages=[{"role": "user", "content": user_prompt}], - tools=[FLASHCARD_TOOL], - tool_choice={"type": "tool", "name": "create_flashcards"} - ) - - # Extract flashcards from tool call and convert to Flashcard objects - if response.content and len(response.content) > 0: - for content_block in response.content: - if content_block.type == "tool_use": - tool_input = cast(Dict[str, Any], content_block.input) - flashcard_dicts = tool_input['flashcards'] - - flashcard_objects = [] - for card in flashcard_dicts: - front_original = card.get('front', '') - back_original = card.get('back', '') - front_processed = process_code_blocks(front_original, CONFIG.syntax_highlighting) - back_processed = process_code_blocks(back_original, CONFIG.syntax_highlighting) - - flashcard = Flashcard( - front=front_processed, - back=back_processed, - note=note, - tags=card.get('tags', note.tags.copy()), - front_original=front_original, - back_original=back_original - ) - flashcard_objects.append(flashcard) + response = self._call_llm( + system_prompt=SYSTEM_PROMPT, + user_prompt=user_prompt, + tools=[FLASHCARD_TOOL], + tool_choice=self._get_tool_choice("create_flashcards") + ) - return flashcard_objects - - console.print("[yellow]WARNING:[/yellow] No flashcards generated - unexpected response format") + if not response: return [] + # Extract flashcards from tool call + try: + message = response.choices[0].message + if hasattr(message, 'tool_calls') and message.tool_calls: + tool_call = message.tool_calls[0] + import json + arguments = json.loads(tool_call.function.arguments) + + flashcard_dicts = arguments['flashcards'] + + flashcard_objects = [] + for card in flashcard_dicts: + front_original = card.get('front', '') + back_original = card.get('back', '') + front_processed = process_code_blocks(front_original, CONFIG.syntax_highlighting) + back_processed = process_code_blocks(back_original, CONFIG.syntax_highlighting) + + flashcard = Flashcard( + front=front_processed, + back=back_processed, + note=note, + tags=card.get('tags', note.tags.copy()), + front_original=front_original, + back_original=back_original + ) + flashcard_objects.append(flashcard) + + return flashcard_objects except Exception as e: - console.print(f"[red]ERROR:[/red] Error generating flashcards: {e}") + console.print(f"[red]ERROR:[/red] Failed to parse flashcards: {e}") return [] - def generate_from_query(self, query: str, target_cards: int, previous_fronts: list = [], deck_examples: list = []) -> List[Flashcard]: - """Generate flashcards based on a user query without source material""" + console.print("[yellow]WARNING:[/yellow] No flashcards generated - unexpected response format") + return [] + def generate_from_query( + self, + query: str, + target_cards: int, + previous_fronts: List[str] = [], + deck_examples: List[Dict[str, str]] = [] + ) -> List[Flashcard]: + """Generate flashcards based on a user query without source material""" card_instruction = self._build_card_instruction(target_cards) dedup_context = self._build_dedup_context(previous_fronts) schema_context = self._build_schema_context(deck_examples) @@ -186,60 +245,58 @@ def generate_from_query(self, query: str, target_cards: int, previous_fronts: li Please {card_instruction} to help someone learn about this topic. Focus on the most important concepts, definitions, and practical information related to this query.{difficulty_context}{dedup_context}{schema_context}""" - try: - response = self.client.messages.create( - model="claude-4-sonnet-20250514", - max_tokens=8000, - system=QUERY_SYSTEM_PROMPT, - messages=[{"role": "user", "content": user_prompt}], - tools=[FLASHCARD_TOOL], - tool_choice={"type": "tool", "name": "create_flashcards"} - ) - - # Extract flashcards from tool call and convert to Flashcard objects - if response.content and len(response.content) > 0: - for content_block in response.content: - if content_block.type == "tool_use": - tool_input = cast(Dict[str, Any], content_block.input) - flashcard_dicts = tool_input.get("flashcards", []) - - # Create virtual Note object for query-based flashcards - virtual_note = Note( - path="query", - filename=f"Query: {query}", - content=query, - tags=["query-generated"], - size=0 - ) - - flashcard_objects = [] - for card in flashcard_dicts: - # Process the front and back content - front_original = card.get('front', '') - back_original = card.get('back', '') - front_processed = process_code_blocks(front_original, CONFIG.syntax_highlighting) - back_processed = process_code_blocks(back_original, CONFIG.syntax_highlighting) + response = self._call_llm( + system_prompt=QUERY_SYSTEM_PROMPT, + user_prompt=user_prompt, + tools=[FLASHCARD_TOOL], + tool_choice=self._get_tool_choice("create_flashcards") + ) - # Create Flashcard object - flashcard = Flashcard( - front=front_processed, - back=back_processed, - note=virtual_note, - tags=card.get('tags', ["query-generated"]), - front_original=front_original, - back_original=back_original - ) - flashcard_objects.append(flashcard) + if not response: + return [] - return flashcard_objects + # Extract flashcards + try: + message = response.choices[0].message + if hasattr(message, 'tool_calls') and message.tool_calls: + tool_call = message.tool_calls[0] + import json + flashcard_dicts = json.loads(tool_call.function.arguments).get("flashcards", []) + + # Create virtual Note object for query-based flashcards + virtual_note = Note( + path="query", + filename=f"Query: {query}", + content=query, + tags=["query-generated"], + size=0 + ) - console.print("[yellow]WARNING:[/yellow] No flashcards generated - unexpected response format") - return [] + flashcard_objects = [] + for card in flashcard_dicts: + front_original = card.get('front', '') + back_original = card.get('back', '') + front_processed = process_code_blocks(front_original, CONFIG.syntax_highlighting) + back_processed = process_code_blocks(back_original, CONFIG.syntax_highlighting) + + flashcard = Flashcard( + front=front_processed, + back=back_processed, + note=virtual_note, + tags=card.get('tags', ["query-generated"]), + front_original=front_original, + back_original=back_original + ) + flashcard_objects.append(flashcard) + return flashcard_objects except Exception as e: - console.print(f"[red]ERROR:[/red] Error generating flashcards from query: {e}") + console.print(f"[red]ERROR:[/red] Failed to parse flashcards: {e}") return [] + console.print("[yellow]WARNING:[/yellow] No flashcards generated - unexpected response format") + return [] + def generate_from_note_query(self, note: Note, query: str, target_cards: int, previous_fronts: List[str] | None = None, deck_examples: List[Dict[str, str]] | None = None) -> List[Flashcard]: """Generate flashcards by extracting specific information from a note based on a query""" if previous_fronts is None: @@ -260,51 +317,49 @@ def generate_from_note_query(self, note: Note, query: str, target_cards: int, pr Please analyze this note and extract information specifically related to the query "{query}". {card_instruction} only for information in the note that directly addresses or relates to this query.""" - try: - response = self.client.messages.create( - model="claude-4-sonnet-20250514", - max_tokens=8000, - system=TARGETED_SYSTEM_PROMPT, - messages=[{"role": "user", "content": user_prompt}], - tools=[FLASHCARD_TOOL], - tool_choice={"type": "tool", "name": "create_flashcards"} - ) + response = self._call_llm( + system_prompt=TARGETED_SYSTEM_PROMPT, + user_prompt=user_prompt, + tools=[FLASHCARD_TOOL], + tool_choice=self._get_tool_choice("create_flashcards") + ) - # Extract flashcards from tool call and convert to Flashcard objects - if response.content and len(response.content) > 0: - for content_block in response.content: - if content_block.type == "tool_use": - tool_input = cast(Dict[str, Any], content_block.input) - flashcard_dicts = tool_input.get("flashcards", []) - - flashcard_objects = [] - for card in flashcard_dicts: - # Process the front and back content - front_original = card.get('front', '') - back_original = card.get('back', '') - front_processed = process_code_blocks(front_original, CONFIG.syntax_highlighting) - back_processed = process_code_blocks(back_original, CONFIG.syntax_highlighting) - - # Create Flashcard object - flashcard = Flashcard( - front=front_processed, - back=back_processed, - note=note, - tags=card.get('tags', note.tags.copy()), - front_original=front_original, - back_original=back_original - ) - flashcard_objects.append(flashcard) - - return flashcard_objects - - console.print("[yellow]WARNING:[/yellow] No flashcards generated - unexpected response format") + if not response: return [] + # Extract flashcards + try: + message = response.choices[0].message + if hasattr(message, 'tool_calls') and message.tool_calls: + tool_call = message.tool_calls[0] + import json + flashcard_dicts = json.loads(tool_call.function.arguments).get("flashcards", []) + + flashcard_objects = [] + for card in flashcard_dicts: + front_original = card.get('front', '') + back_original = card.get('back', '') + front_processed = process_code_blocks(front_original, CONFIG.syntax_highlighting) + back_processed = process_code_blocks(back_original, CONFIG.syntax_highlighting) + + flashcard = Flashcard( + front=front_processed, + back=back_processed, + note=note, + tags=card.get('tags', note.tags.copy()), + front_original=front_original, + back_original=back_original + ) + flashcard_objects.append(flashcard) + + return flashcard_objects except Exception as e: - console.print(f"[red]ERROR:[/red] Error generating targeted flashcards: {e}") + console.print(f"[red]ERROR:[/red] Failed to parse flashcards: {e}") return [] + console.print("[yellow]WARNING:[/yellow] No flashcards generated - unexpected response format") + return [] + def find_with_agent(self, natural_request: str, sample_size: int | None = None, bias_strength: float | None = None) -> List[Note]: """Use multi-turn agent with tool calling to find notes via iterative DQL refinement""" from datetime import datetime @@ -321,70 +376,71 @@ def find_with_agent(self, natural_request: str, sample_size: int | None = None, Find the most relevant notes for this request using DQL queries. Start with an initial query, analyze the results, and refine as needed.""" # Multi-turn conversation with tool calling - messages: List[MessageParam] = [{"role": "user", "content": user_prompt}] + messages = [ + {"role": "system", "content": MULTI_TURN_DQL_AGENT_PROMPT}, + {"role": "user", "content": user_prompt} + ] max_turns = 8 selected_notes = [] - last_results = [] # Keep track of last query results - all_results = {} # Accumulate all results by path for validation - has_dql_results = False # Track if we've gotten at least one DQL result + last_results = [] + all_results = {} + has_dql_results = False - for _ in range(max_turns): + for turn in range(max_turns): try: - tool_choice_param: ToolChoiceParam + # Determine available tools if not has_dql_results: available_tools = [DQL_EXECUTION_TOOL] - tool_choice_param = {"type": "tool", "name": "execute_dql_query"} + tool_choice = self._get_tool_choice("execute_dql_query") else: available_tools = [DQL_EXECUTION_TOOL, FINALIZE_SELECTION_TOOL] - tool_choice_param = {"type": "any"} + tool_choice = "auto" - response = self.client.messages.create( - model="claude-4-sonnet-20250514", - max_tokens=3000, - system=MULTI_TURN_DQL_AGENT_PROMPT, + response = cast(ModelResponse, completion( + model=self.model, messages=messages, tools=available_tools, - tool_choice=tool_choice_param - ) + tool_choice=tool_choice, + max_tokens=3000 + )) - messages.append({"role": "assistant", "content": response.content}) + message = response.choices[0].message + messages.append({"role": "assistant", "content": message.content or "", "tool_calls": message.tool_calls if hasattr(message, 'tool_calls') else None}) tool_results = [] final_selection = None - for content_block in response.content: - if content_block.type == "tool_use": - tool_name = content_block.name - tool_input = cast(Dict[str, Any], content_block.input) + if hasattr(message, 'tool_calls') and message.tool_calls: + for tool_call in message.tool_calls: + tool_name = tool_call.function.name + import json + tool_input = json.loads(tool_call.function.arguments) if tool_name == "execute_dql_query": dql_query = tool_input["query"] - reasoning = tool_input["reasoning"] + reasoning = tool_input.get("reasoning", "") console.print(f"[cyan]Agent:[/cyan] {reasoning}") console.print(f"[dim]Query:[/dim] {dql_query}") try: - # Execute the DQL query from obsidianki.cli.services import OBSIDIAN results = OBSIDIAN.dql(dql_query) if results is None: results = [] - # Apply filtering (folders, excluded tags) + # Apply filtering filtered_results = [] for result in results: note_path = result.path note_tags = result.tags or [] - # Apply search_folders filtering if CONFIG.search_folders: path_matches = any(note_path.startswith(f"{folder}/") for folder in CONFIG.search_folders) if not path_matches: continue - # Apply excluded tags filtering excluded_tags = CONFIG.get_excluded_tags() if excluded_tags and any(tag in note_tags for tag in excluded_tags): continue @@ -394,48 +450,33 @@ def find_with_agent(self, natural_request: str, sample_size: int | None = None, results = filtered_results console.print(f"[cyan]Agent:[/cyan] Found {len(results)} notes") - last_results = results # Store for potential auto-finalization - has_dql_results = True # Mark that we now have DQL results + last_results = results + has_dql_results = True - # Accumulate all results by path for validation for result in results: - # Handle Note objects directly - if hasattr(result, 'path'): - path = result.path - else: - # Fallback for dict format - path = result.get('result', {}).get('path') + path = result.path if hasattr(result, 'path') else result.get('result', {}).get('path') if path: all_results[path] = result - # Prepare result summary for AI + # Prepare result summary if len(results) == 0: result_summary = "No notes found matching this query." elif len(results) <= AI_RESULT_SET_SIZE: - # Show detailed results for small result sets result_list = [] for i, result in enumerate(results[:AI_RESULT_SET_SIZE]): - # Handle Note objects directly - if hasattr(result, 'path'): - path = result.path - name = result.filename - tags = result.tags - size = result.size - else: - # Fallback for dict format - note = result.get('result', {}) - path = note.get('path', 'Unknown') - name = note.get('name', 'Unknown') - tags = note.get('tags', []) - size = note.get('size', 0) + path = result.path if hasattr(result, 'path') else result.get('result', {}).get('path', 'Unknown') + name = result.filename if hasattr(result, 'filename') else result.get('result', {}).get('name', 'Unknown') + tags = result.tags if hasattr(result, 'tags') else result.get('result', {}).get('tags', []) + size = result.size if hasattr(result, 'size') else result.get('result', {}).get('size', 0) result_list.append(f"{i+1}. {name} ({path}) - {size} chars, tags: {tags}") result_summary = f"Found {len(results)} notes:\n" + "\n".join(result_list) else: - # Show summary for large result sets result_summary = f"Found {len(results)} notes - this may be too many. Consider refining your query to be more specific." tool_results.append({ - "tool_use_id": content_block.id, + "tool_call_id": tool_call.id, + "role": "tool", + "name": tool_name, "content": result_summary }) @@ -443,13 +484,15 @@ def find_with_agent(self, natural_request: str, sample_size: int | None = None, error_msg = f"DQL Error: {str(e)}" console.print(f"[yellow]{error_msg}[/yellow]") tool_results.append({ - "tool_use_id": content_block.id, + "tool_call_id": tool_call.id, + "role": "tool", + "name": tool_name, "content": error_msg }) elif tool_name == "finalize_note_selection": selected_paths = tool_input["selected_paths"] - reasoning = tool_input["reasoning"] + reasoning = tool_input.get("reasoning", "") console.print(f"[cyan]Agent:[/cyan] {reasoning}") console.print(f"[cyan]Agent:[/cyan] Selected {len(selected_paths)} notes for processing") @@ -467,23 +510,15 @@ def find_with_agent(self, natural_request: str, sample_size: int | None = None, console.print(f"[cyan]Agent:[/cyan] Proceeding with {len(final_selection)} valid selections") tool_results.append({ - "tool_use_id": content_block.id, + "tool_call_id": tool_call.id, + "role": "tool", + "name": tool_name, "content": f"Selection finalized: {len(final_selection)} notes will be processed." }) # Add tool results to conversation if tool_results: - for tool_result in tool_results: - messages.append({ - "role": "user", - "content": [ - { - "type": "tool_result", - "tool_use_id": tool_result["tool_use_id"], - "content": tool_result["content"] - } - ] - }) + messages.extend(tool_results) # If agent finalized selection, we're done if final_selection is not None: @@ -494,60 +529,19 @@ def find_with_agent(self, natural_request: str, sample_size: int | None = None, console.print(f"[red]ERROR:[/red] Agent conversation failed: {e}") return [] - if not selected_notes: - # Force agent to finalize selection if it hasn't already - if last_results: - console.print(f"[cyan]Agent:[/cyan] Forcing finalization of {len(last_results)} available notes") - - try: - # Send final request forcing finalize_note_selection - response = self.client.messages.create( - model="claude-4-sonnet-20250514", - max_tokens=3000, - system=MULTI_TURN_DQL_AGENT_PROMPT, - messages=messages + [{"role": "user", "content": "Please finalize your note selection now using the finalize_note_selection tool."}], - tools=[FINALIZE_SELECTION_TOOL], - tool_choice={"type": "tool", "name": "finalize_note_selection"} - ) - - # Process the forced finalization - for content_block in response.content: - if content_block.type == "tool_use" and content_block.name == "finalize_note_selection": - tool_input = cast(Dict[str, Any], content_block.input) - selected_paths = tool_input["selected_paths"] - reasoning = tool_input["reasoning"] - - console.print(f"[cyan]Agent:[/cyan] {reasoning}") - console.print(f"[cyan]Agent:[/cyan] Selected {len(selected_paths)} notes for processing") - - # Find the corresponding note objects from all accumulated results - final_selection = [] - missing_paths = [] - for path in selected_paths: - if path in all_results: - final_selection.append(all_results[path]) - else: - missing_paths.append(path) - - # Warn about any missing paths - if missing_paths: - console.print(f"[yellow]Warning:[/yellow] Agent selected {len(missing_paths)} paths not found in query results: {missing_paths}") - console.print(f"[cyan]Agent:[/cyan] Proceeding with {len(final_selection)} valid selections") + # Force finalization if needed + if not selected_notes and last_results: + console.print(f"[cyan]Agent:[/cyan] Forcing finalization of {len(last_results)} available notes") + selected_notes = last_results - selected_notes = final_selection - break - - except Exception as e: - console.print(f"[red]ERROR:[/red] Failed to force finalization: {e}") - return [] - - if not selected_notes: - console.print("[yellow]Agent could not finalize a selection[/yellow]") - return [] + if not selected_notes: + console.print("[yellow]Agent could not finalize a selection[/yellow]") + return [] - # Apply weighted sampling to final selection if needed + # Apply sampling if needed target_count = sample_size if sample_size else len(selected_notes) if target_count < len(selected_notes): + from obsidianki.cli.services import OBSIDIAN bias = bias_strength if bias_strength is not None else 1.0 sampled_notes = OBSIDIAN._weighted_sample(selected_notes, target_count, bias) else: @@ -561,7 +555,7 @@ def edit_cards(self, cards: List[Dict[str, str]], query: str) -> List[Dict[str, if not cards: return [] - # Build card context using original text (strip HTML for cleaner AI input) + # Build card context cards_context = "" for i, card in enumerate(cards, 1): front_clean = strip_html(card['front']) @@ -593,55 +587,52 @@ def edit_cards(self, cards: List[Dict[str, str]], query: str) -> List[Dict[str, - Use markdown syntax with triple backticks for code blocks (```language\\ncode\\n```) - Do NOT use HTML tags like
, , 
, etc.""" + response = self._call_llm( + system_prompt=edit_system_prompt, + user_prompt=edit_prompt, + tools=[FLASHCARD_TOOL], + tool_choice=self._get_tool_choice("create_flashcards"), + max_tokens=4000 + ) + + if not response: + return cards + try: - response = self.client.messages.create( - model="claude-4-sonnet-20250514", - max_tokens=4000, - system=edit_system_prompt, - messages=[ - {"role": "user", "content": edit_prompt} - ], - tools=[FLASHCARD_TOOL], - tool_choice={"type": "tool", "name": "create_flashcards"} - ) + message = response.choices[0].message + if hasattr(message, 'tool_calls') and message.tool_calls: + tool_call = message.tool_calls[0] + import json + flashcard_data = json.loads(tool_call.function.arguments) + + if "flashcards" in flashcard_data: + edited_cards = [] + for flashcard in flashcard_data["flashcards"]: + if "front" in flashcard and "back" in flashcard: + front_original = flashcard["front"] + back_original = flashcard["back"] - if not response.content: - console.print("[yellow]WARNING:[/yellow] No response from AI for card editing") - return cards - - edited_cards = [] - - for content_block in response.content: - if content_block.type == "tool_use" and content_block.name == "create_flashcards": - tool_input = cast(Dict[str, Any], content_block.input) - if "flashcards" in tool_input: - for flashcard_data in tool_input["flashcards"]: - if "front" in flashcard_data and "back" in flashcard_data: - # Store original text before processing - front_original = flashcard_data["front"] - back_original = flashcard_data["back"] - - # Process code blocks like other flashcard generation - front_processed = process_code_blocks(front_original, CONFIG.syntax_highlighting) - back_processed = process_code_blocks(back_original, CONFIG.syntax_highlighting) - - edited_cards.append({ - "front": front_processed, - "back": back_processed, - "front_original": front_original, - "back_original": back_original, - "origin": flashcard_data.get("origin", "") - }) + front_processed = process_code_blocks(front_original, CONFIG.syntax_highlighting) + back_processed = process_code_blocks(back_original, CONFIG.syntax_highlighting) - if len(edited_cards) != len(cards): - console.print(f"[yellow]WARNING:[/yellow] Expected {len(cards)} edited cards, got {len(edited_cards)}.") - console.print(f"[yellow]AI returned incomplete results. Using original cards.[/yellow]") - return cards + edited_cards.append({ + "front": front_processed, + "back": back_processed, + "front_original": front_original, + "back_original": back_original, + "origin": flashcard.get("origin", "") + }) - return edited_cards + if len(edited_cards) != len(cards): + console.print(f"[yellow]WARNING:[/yellow] Expected {len(cards)} edited cards, got {len(edited_cards)}.") + console.print(f"[yellow]AI returned incomplete results. Using original cards.[/yellow]") + return cards + return edited_cards except Exception as e: import traceback console.print(f"[red]ERROR:[/red] Failed to edit cards: {e}") console.print(f"[dim]{traceback.format_exc()}[/dim]") - return cards \ No newline at end of file + return cards + + return cards diff --git a/obsidianki/ai/models.py b/obsidianki/ai/models.py new file mode 100644 index 0000000..f229178 --- /dev/null +++ b/obsidianki/ai/models.py @@ -0,0 +1,50 @@ +MODEL_MAP = { + "Claude Sonnet 4.5": { + "provider": "anthropic", + "url": "https://console.anthropic.com/", + "model": "claude-sonnet-4-5", + "key_name": "ANTHROPIC_API_KEY" + }, + "Claude Opus 4": { + "provider": "anthropic", + "model": "claude-opus-4-1", + "key_name": "ANTHROPIC_API_KEY", + "url": "https://console.anthropic.com/" + }, + "GPT-5": { + "provider": "openai", + "model": "gpt-5", + "key_name": "OPENAI_API_KEY", + "url": "https://platform.openai.com/api-keys" + }, + "Gemini 3": { + "provider": "google", + "model": "gemini/gemini-2.5-pro", + "key_name": "GEMINI_API_KEY", + "url": "https://makersuite.google.com/app/apikey" + }, + "GPT-4o": { + "provider": "openai", + "model": "gpt-4o", + "key_name": "OPENAI_API_KEY", + "url": "https://platform.openai.com/api-keys" + }, + "GPT-4o Mini": { + "provider": "openai", + "model": "gpt-4o-mini", + "key_name": "OPENAI_API_KEY", + "url": "https://platform.openai.com/api-keys" + }, + "Gemini 2.5 Flash": { + "provider": "google", + "model": "gemini/gemini-2.5-flash", + "key_name": "GEMINI_API_KEY", + "url": "https://makersuite.google.com/app/apikey" + }, + "DeepSeek V3.1": { + "provider": "deepseek", + "model": "deepseek/deepseek-chat", + "key_name": "DEEPSEEK_API_KEY", + "url": "https://console.deepseek.com/" + } +} \ No newline at end of file diff --git a/obsidianki/ai/tools.py b/obsidianki/ai/tools.py index 9f89dc7..1996820 100644 --- a/obsidianki/ai/tools.py +++ b/obsidianki/ai/tools.py @@ -1,71 +1,78 @@ -from anthropic.types import ToolParam - -FLASHCARD_TOOL: ToolParam = { - "name": "create_flashcards", - "description": "Create flashcards from note content with front (question) and back (answer)", - "input_schema": { - "type": "object", - "properties": { - "flashcards": { - "type": "array", - "description": "Array of flashcards extracted from the note", - "items": { - "type": "object", - "properties": { - "front": { - "type": "string", - "description": "The question or prompt for the flashcard" +FLASHCARD_TOOL: dict = { + "type": "function", + "function": { + "name": "create_flashcards", + "description": "Create flashcards from note content with front (question) and back (answer)", + "parameters": { + "type": "object", + "properties": { + "flashcards": { + "type": "array", + "description": "Array of flashcards extracted from the note", + "items": { + "type": "object", + "properties": { + "front": { + "type": "string", + "description": "The question or prompt for the flashcard" + }, + "back": { + "type": "string", + "description": "The answer or information for the flashcard" + } }, - "back": { - "type": "string", - "description": "The answer or information for the flashcard" - } - }, - "required": ["front", "back"] + "required": ["front", "back"] + } } - } - }, - "required": ["flashcards"] + }, + "required": ["flashcards"] + } } } # DQL Execution Tool for multi-turn agent -DQL_EXECUTION_TOOL: ToolParam = { - "name": "execute_dql_query", - "description": "Execute a DQL query against the Obsidian vault and get results", - "input_schema": { - "type": "object", - "properties": { - "query": { - "type": "string", - "description": "The DQL query to execute" +DQL_EXECUTION_TOOL: dict = { + "type": "function", + "function": { + "name": "execute_dql_query", + "description": "Execute a DQL query against the Obsidian vault and get results", + "parameters": { + "type": "object", + "properties": { + "query": { + "type": "string", + "description": "The DQL query to execute" + }, + "reasoning": { + "type": "string", + "description": "Brief explanation of what this query is trying to find" + } }, - "reasoning": { - "type": "string", - "description": "Brief explanation of what this query is trying to find" - } - }, - "required": ["query", "reasoning"] + "required": ["query", "reasoning"] + } } } # Final selection tool for multi-turn agent -FINALIZE_SELECTION_TOOL: ToolParam = { - "name": "finalize_note_selection", - "description": "Finalize the selection of notes that best match the user's request", - "input_schema": { - "type": "object", - "properties": { - "selected_paths": { - "type": "array", - "items": {"type": "string"}, - "description": "Array of note paths to process for flashcard generation" +FINALIZE_SELECTION_TOOL: dict = { + "type": "function", + "function": { + "name": "finalize_note_selection", + "description": "Finalize the selection of notes that best match the user's request", + "parameters": { + "type": "object", + "properties": { + "selected_paths": { + "type": "array", + "items": {"type": "string"}, + "description": "Array of note paths to process for flashcard generation" + }, + "reasoning": { + "type": "string", + "description": "Brief explanation of why these notes were selected" + } }, - "reasoning": { - "type": "string", - "description": "Brief explanation of why these notes were selected" - } - }, - "required": ["selected_paths", "reasoning"] + "required": ["selected_paths", "reasoning"] + } } } \ No newline at end of file diff --git a/obsidianki/cli/commands/config_cmd.py b/obsidianki/cli/commands/config_cmd.py index 8bb6018..9c40489 100644 --- a/obsidianki/cli/commands/config_cmd.py +++ b/obsidianki/cli/commands/config_cmd.py @@ -16,6 +16,7 @@ def handle_config_command(args): "config": "List all configuration settings", "config get ": "Get a configuration value", "config set ": "Set a configuration value", + "config set model \"\"": "Set model", "config reset": "Reset configuration to defaults", "config where": "Show configuration directory path" }) @@ -68,10 +69,30 @@ def handle_config_command(args): from obsidianki.cli.config import DEFAULT_CONFIG key_upper = args.key.upper() + + # Special handling for "model" - allows human-friendly names + if key_upper == 'MODEL': + from obsidianki.ai.models import MODEL_MAP + + if args.value in MODEL_MAP: + user_config["MODEL"] = args.value + + with open(CONFIG_FILE, 'w') as f: + json.dump(user_config, f, indent=2) + + console.print(f"[green]✓[/green] Set model to [bold]{args.value}[/bold]") + return + else: + console.print(f"[red]Invalid model: {args.value}[/red]") + console.print("[dim]Valid options:[/dim]") + for model_name in MODEL_MAP.keys(): + console.print(f" - {model_name}") + return + # Check if key exists in DEFAULT_CONFIG (support new config keys) if key_upper not in DEFAULT_CONFIG: console.print(f"[red]Configuration key '{args.key}' not found.[/red]") - console.print("[dim]Use 'oki config list' to see available keys.[/dim]") + console.print("[dim]Use 'oki config' to see available keys.[/dim]") return # Try to convert value to appropriate type @@ -83,7 +104,7 @@ def handle_config_command(args): if key_upper == 'DIFFICULTY': if value not in ('easy', 'normal', 'hard', 'none'): console.print(f"[red]Invalid difficulty: {value}[/red]") - console.print("[dim]Valid options: easy, normal, hard[/dim]") + console.print("[dim]Valid options: easy, normal, hard, none[/dim]") return if isinstance(current_value, bool): diff --git a/obsidianki/cli/commands/deck_cmd.py b/obsidianki/cli/commands/deck_cmd.py index 81271c4..96fe26a 100644 --- a/obsidianki/cli/commands/deck_cmd.py +++ b/obsidianki/cli/commands/deck_cmd.py @@ -4,7 +4,6 @@ from rich.markup import escape from rich.panel import Panel -from obsidianki.cli.services import ANKI from obsidianki.cli.config import console from obsidianki.cli.utils import strip_html from obsidianki.cli.help_utils import show_simple_help @@ -12,6 +11,7 @@ def handle_deck_command(args): """Handle deck management commands""" + from obsidianki.cli.services import ANKI # Handle help request if args.help: @@ -23,17 +23,15 @@ def handle_deck_command(args): }) return - anki = ANKI - # Test connection first - if not anki.test_connection(): + if not ANKI.test_connection(): console.print("[red]ERROR:[/red] Cannot connect to AnkiConnect") console.print("[dim]Make sure Anki is running with AnkiConnect add-on installed[/dim]") return if args.deck_action is None: # Default action: list decks - deck_names = anki.get_decks() + deck_names = ANKI.get_decks() if not deck_names: console.print("[yellow]No decks found[/yellow]") @@ -49,7 +47,7 @@ def handle_deck_command(args): console.print(f"[dim]Found {len(deck_names)} decks:[/dim]") console.print() for deck_name in sorted(deck_names): - stats = anki.get_stats(deck_name) + stats = ANKI.get_stats(deck_name) total_cards = stats.get("total_cards", 0) console.print(f" [cyan]{deck_name}[/cyan]") @@ -69,7 +67,7 @@ def handle_deck_command(args): console.print(f"[cyan]Renaming deck:[/cyan] [bold]{old_name}[/bold] → [bold]{new_name}[/bold]") - if anki.rename_deck(old_name, new_name): + if ANKI.rename_deck(old_name, new_name): console.print(f"[green]✓[/green] Successfully renamed deck to '[cyan]{new_name}[/cyan]'") else: console.print("[red]Failed to rename deck[/red]") @@ -82,7 +80,7 @@ def handle_deck_command(args): limit = args.limit # Check if deck exists - deck_names = anki.get_decks() + deck_names = ANKI.get_decks() if deck_name not in deck_names: console.print(f"[red]ERROR:[/red] Deck '[cyan]{deck_name}[/cyan]' not found") console.print("\n[dim]Available decks:[/dim]") @@ -95,7 +93,7 @@ def handle_deck_command(args): console.print() # Search for cards - results = anki.search_cards(deck_name, query, limit) + results = ANKI.search_cards(deck_name, query, limit) if not results: console.print(f"[yellow]No cards found matching '{query}'[/yellow]") diff --git a/obsidianki/cli/config.py b/obsidianki/cli/config.py index 0866cf3..fe18eba 100644 --- a/obsidianki/cli/config.py +++ b/obsidianki/cli/config.py @@ -37,7 +37,8 @@ "SYNTAX_HIGHLIGHTING": True, # Enable syntax highlighting for code blocks in flashcards "UPFRONT_BATCHING": False, # Process all notes in parallel instead of one-by-one "BATCH_SIZE_LIMIT": 20, # Maximum notes to process in batch mode - "BATCH_CARD_LIMIT": 100 # Maximum total cards in batch mode + "BATCH_CARD_LIMIT": 100, # Maximum total cards in batch mode + "MODEL": "Claude Sonnet 4.5" # AI model to use (Claude Sonnet 4, GPT-5, Gemini 3 Pro Preview, etc.) } class Config: diff --git a/obsidianki/cli/interactive/edit_mode.py b/obsidianki/cli/interactive/edit_mode.py index 8181fe1..5aa2d25 100644 --- a/obsidianki/cli/interactive/edit_mode.py +++ b/obsidianki/cli/interactive/edit_mode.py @@ -4,7 +4,6 @@ from rich.prompt import Prompt from obsidianki.cli.models import Note, Flashcard -from obsidianki.cli.services import ANKI, AI from obsidianki.cli.config import CONFIG, console from obsidianki.cli.interactive.approval import approve_flashcard from obsidianki.cli.interactive.card_selector import create_card_selector @@ -14,6 +13,8 @@ def edit_mode(args): """ Entry point for interactive editing of existing flashcards. """ + from obsidianki.cli.services import ANKI, AI + deck_name = args.deck if args.deck else CONFIG.deck console.print(Panel("ObsidianKi - Editing mode", style="bold blue")) diff --git a/obsidianki/cli/wizard.py b/obsidianki/cli/wizard.py index d4b2ad2..d131628 100644 --- a/obsidianki/cli/wizard.py +++ b/obsidianki/cli/wizard.py @@ -7,9 +7,12 @@ def setup(force_full_setup=False): """Interactive setup to configure API keys and preferences""" + import questionary + console.print(Panel(Text("ObsidianKi Setup", style="bold blue"), style="blue")) step_num = 1 + model_choice = None # Initialize to None, will be set if API keys are configured CONFIG_DIR.mkdir(parents=True, exist_ok=True) @@ -22,14 +25,27 @@ def setup(force_full_setup=False): console.print("[red]ERROR:[/red] Obsidian API key is required. Setup aborted.") return - console.print("\n Get Anthropic API key from: [blue]https://console.anthropic.com/[/blue]") - anthropic_key = Prompt.ask(" Enter your Anthropic API key", password=True).strip() - if not anthropic_key: - console.print("[red]ERROR:[/red] Anthropic API key is required. Setup aborted.") + console.print("\n [cyan]AI Model Selection[/cyan]") + + from obsidianki.ai.models import MODEL_MAP + model_choice = questionary.select( + " Select model:", + choices=list(MODEL_MAP.keys()), + default="Claude Sonnet 4.5", + instruction="" + ).ask() + + model_info = MODEL_MAP[model_choice] + + console.print(f"\n Get API key from: [blue]{model_info['url']}[/blue]") + + ai_key = Prompt.ask(f" Enter your API key", password=True).strip() + if not ai_key: + console.print(f"[red]ERROR:[/red] API key is required. Setup aborted.") return env_content = f"""OBSIDIAN_API_KEY={obsidian_key} -ANTHROPIC_API_KEY={anthropic_key} +{model_info['key_name']}={ai_key} """ try: @@ -52,17 +68,19 @@ def setup(force_full_setup=False): notes_to_sample = IntPrompt.ask(" How many notes to sample?", default=CONFIG.notes_to_sample) days_old = IntPrompt.ask(" Only process notes older than X days?", default=CONFIG.days_old) - sampling_mode = Prompt.ask( - " Sampling mode", + sampling_mode = questionary.select( + " Sampling mode:", choices=["random", "weighted"], - default=CONFIG.sampling_mode - ) + default=CONFIG.sampling_mode, + instruction="" + ).ask() - card_type = Prompt.ask( - " Card type", + card_type = questionary.select( + " Card type:", choices=["basic", "custom"], - default=CONFIG.card_type - ) + default=CONFIG.card_type, + instruction="" + ).ask() console.print("\n [cyan]Approval Settings[/cyan]") approve_notes = Confirm.ask( @@ -89,7 +107,7 @@ def setup(force_full_setup=False): from obsidianki.cli.config import DEFAULT_CONFIG user_config = DEFAULT_CONFIG.copy() - user_config.update({ + config_update = { "MAX_CARDS": max_cards, "NOTES_TO_SAMPLE": notes_to_sample, "DAYS_OLD": days_old, @@ -99,7 +117,11 @@ def setup(force_full_setup=False): "APPROVE_CARDS": approve_cards, "DEDUPLICATE_VIA_HISTORY": deduplicate_via_history, "SYNTAX_HIGHLIGHTING": syntax_highlighting, - }) + } + # Only update MODEL if it was set during this setup run + if model_choice is not None: + config_update["MODEL"] = model_choice + user_config.update(config_update) try: CONFIG.save(user_config) diff --git a/obsidianki/main.py b/obsidianki/main.py index c839e89..3012778 100644 --- a/obsidianki/main.py +++ b/obsidianki/main.py @@ -6,7 +6,7 @@ def _excepthook(exc_type, exc_value, exc_traceback): if exc_type is KeyboardInterrupt: sys.exit(130) else: - print(f"\nERROR: {exc_value}", file=sys.stderr) + console.print(f"[red]ERROR:[/red] {exc_value}") sys.exit(1) sys.excepthook = _excepthook @@ -14,13 +14,6 @@ def _excepthook(exc_type, exc_value, exc_traceback): from rich.text import Text from obsidianki.cli.config import console, ENV_FILE, CONFIG_FILE -from obsidianki.cli.commands.config_cmd import handle_config_command -from obsidianki.cli.commands.tag_cmd import handle_tag_command -from obsidianki.cli.commands.history_cmd import handle_history_command -from obsidianki.cli.commands.deck_cmd import handle_deck_command -from obsidianki.cli.commands.template_cmd import handle_template_command -from obsidianki.cli.commands.hide_cmd import handle_hide_command -from obsidianki.cli.interactive.edit_mode import edit_mode def show_main_help(): """Display the main help screen""" @@ -201,24 +194,31 @@ def main(): return 0 if args.command == 'config': + from obsidianki.cli.commands.config_cmd import handle_config_command handle_config_command(args) return 0 elif args.command == 'history': + from obsidianki.cli.commands.history_cmd import handle_history_command handle_history_command(args) return 0 elif args.command in ['tag', 'tags']: + from obsidianki.cli.commands.tag_cmd import handle_tag_command handle_tag_command(args) return 0 elif args.command == 'deck': + from obsidianki.cli.commands.deck_cmd import handle_deck_command handle_deck_command(args) return 0 elif args.command in ['template', 'templates']: + from obsidianki.cli.commands.template_cmd import handle_template_command handle_template_command(args) return 0 elif args.command == 'hide': + from obsidianki.cli.commands.hide_cmd import handle_hide_command handle_hide_command(args) return 0 elif args.command == 'edit': + from obsidianki.cli.interactive.edit_mode import edit_mode edit_mode(args) return 0 diff --git a/pyproject.toml b/pyproject.toml index bdd86ac..beff2d4 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -15,10 +15,11 @@ classifiers = [] dependencies = [ "requests>=2.25.0", "python-dotenv>=0.19.0", - "anthropic>=0.3.0", + "litellm>=1.8.0", "rich>=13.0.0", "urllib3>=1.26.0", - "pygments>=2.10.0" + "pygments>=2.10.0", + "questionary>=2.0.0" ] [project.optional-dependencies] diff --git a/tests/test_multi_provider.py b/tests/test_multi_provider.py new file mode 100644 index 0000000..ad84ab6 --- /dev/null +++ b/tests/test_multi_provider.py @@ -0,0 +1,156 @@ +"""Tests for multi-provider LLM support via LiteLLM""" +import pytest +import os +from unittest.mock import patch + +from obsidianki.ai.models import MODEL_MAP + + +class TestModelMap: + """Test the MODEL_MAP configuration""" + + def test_model_map_has_expected_models(self): + """Verify MODEL_MAP contains all expected models""" + expected_models = [ + "Claude Sonnet 4.5", + "Claude Opus 4", + "GPT-5", + "Gemini 3", + "GPT-4o", + "GPT-4o Mini", + "Gemini 2.5 Flash", + "DeepSeek V3.1" + ] + + for model in expected_models: + assert model in MODEL_MAP, f"Model '{model}' not found in MODEL_MAP" + + def test_model_map_entries_have_required_fields(self): + """Verify each MODEL_MAP entry has provider, model, and key_name""" + required_fields = ['provider', 'model', 'key_name'] + + for model_name, model_info in MODEL_MAP.items(): + for field in required_fields: + assert field in model_info, f"Model '{model_name}' missing field '{field}'" + assert model_info[field], f"Model '{model_name}' has empty '{field}'" + + def test_anthropic_models_use_correct_provider(self): + """Verify Anthropic models use 'anthropic' provider""" + claude_models = ["Claude Sonnet 4.5", "Claude Opus 4"] + + for model in claude_models: + assert MODEL_MAP[model]["provider"] == "anthropic" + assert MODEL_MAP[model]["key_name"] == "ANTHROPIC_API_KEY" + + def test_openai_models_use_correct_provider(self): + """Verify OpenAI models use 'openai' provider""" + openai_models = ["GPT-5", "GPT-4o", "GPT-4o Mini"] + + for model in openai_models: + assert MODEL_MAP[model]["provider"] == "openai" + assert MODEL_MAP[model]["key_name"] == "OPENAI_API_KEY" + + def test_google_models_use_correct_provider(self): + """Verify Google models use 'google' provider""" + google_models = ["Gemini 3", "Gemini 2.5 Flash"] + + for model in google_models: + assert MODEL_MAP[model]["provider"] == "google" + assert MODEL_MAP[model]["key_name"] == "GEMINI_API_KEY" + + def test_deepseek_models_use_correct_provider(self): + """Verify DeepSeek models use 'deepseek' provider""" + assert MODEL_MAP["DeepSeek V3.1"]["provider"] == "deepseek" + assert MODEL_MAP["DeepSeek V3.1"]["key_name"] == "DEEPSEEK_API_KEY" + + +class TestFlashcardAIModelSelection: + """Test FlashcardAI model initialization with different providers""" + + def test_ai_client_uses_claude_by_default(self): + """Test that FlashcardAI defaults to Claude Sonnet 4.5""" + with patch.dict(os.environ, {'ANTHROPIC_API_KEY': 'test_key'}): + from obsidianki.ai.client import FlashcardAI + + # Mock CONFIG to not have model set + with patch('obsidianki.ai.client.CONFIG') as mock_config: + mock_config.model = 'Claude Sonnet 4.5' + + ai = FlashcardAI() + + assert ai.provider == "anthropic" + assert "claude" in ai.model.lower() + + def test_ai_client_respects_model_config(self): + """Test that FlashcardAI uses model from CONFIG""" + test_cases = [ + ("GPT-5", "openai", "gpt-5"), + ("Claude Opus 4", "anthropic", "claude-opus-4-1"), + ("Gemini 3", "google", "gemini/gemini-2.5-pro"), + ] + + for model_name, expected_provider, expected_model in test_cases: + model_info = MODEL_MAP[model_name] + api_key_name = model_info["key_name"] + + with patch.dict(os.environ, {api_key_name: 'test_key'}): + from obsidianki.ai.client import FlashcardAI + + with patch('obsidianki.ai.client.CONFIG') as mock_config: + mock_config.model = model_name + + ai = FlashcardAI() + + assert ai.provider == expected_provider, \ + f"Model {model_name} should use provider {expected_provider}" + assert ai.model == expected_model, \ + f"Model {model_name} should map to {expected_model}" + + +class TestModelConfiguration: + """Test model configuration via config command""" + + def test_config_accepts_valid_model_names(self): + """Test that config command accepts valid model names""" + from obsidianki.cli.config import CONFIG + + valid_models = ["GPT-5", "Claude Sonnet 4.5", "Gemini 2.5 Flash"] + + for model in valid_models: + assert model in MODEL_MAP, \ + f"Test assumes {model} is in MODEL_MAP but it's not" + + def test_all_model_map_keys_are_user_friendly(self): + """Verify MODEL_MAP keys are human-friendly, not technical IDs""" + for model_name in MODEL_MAP.keys(): + # Should not be technical model IDs like "claude-sonnet-4.5-20250514" + assert not model_name.startswith("claude-"), \ + f"Model key '{model_name}' should be human-friendly, not technical ID" + assert not model_name.startswith("gpt-"), \ + f"Model key '{model_name}' should be human-friendly, not technical ID" + assert not model_name.startswith("gemini-"), \ + f"Model key '{model_name}' should be human-friendly, not technical ID" + + # Should contain spaces or be a proper name + assert " " in model_name or model_name[0].isupper(), \ + f"Model key '{model_name}' should be human-friendly with spaces or proper capitalization" + + +class TestBackwardsCompatibility: + """Test backwards compatibility with existing configs""" + + def test_anthropic_api_key_still_works(self): + """Verify ANTHROPIC_API_KEY environment variable still works""" + with patch.dict(os.environ, {'ANTHROPIC_API_KEY': 'sk-ant-test123'}): + from obsidianki.ai.client import FlashcardAI + + with patch('obsidianki.ai.client.CONFIG') as mock_config: + mock_config.model = 'Claude Sonnet 4.5' + + # Should not raise an error about missing API key + ai = FlashcardAI() + assert ai.provider == "anthropic" + + +if __name__ == "__main__": + pytest.main([__file__, "-v"]) diff --git a/tests/test_setup.py b/tests/test_setup.py index 957c40f..e7e5cc9 100644 --- a/tests/test_setup.py +++ b/tests/test_setup.py @@ -31,7 +31,7 @@ def mock_prompt_ask(prompt_text, default=None, choices=None, password=False, **k prompt_lower = prompt_text.lower() # API keys - if 'anthropic' in prompt_lower: + if 'anthropic' in prompt_lower or 'api key' in prompt_lower: return mock_responses['api_key'] elif 'obsidian' in prompt_lower: return mock_responses['obsidian_key'] @@ -75,9 +75,30 @@ def mock_confirm_ask(prompt_text, default=None, **kwargs): return default if default is not None else True + def mock_questionary_select(message, choices=None, default=None, **kwargs): + """Mock questionary.select() - returns object with .ask() method""" + class MockResponse: + def __init__(self, value): + self.value = value + def ask(self): + return self.value + + message_lower = message.lower() + + if 'model' in message_lower: + return MockResponse("Claude Sonnet 4.5") + elif 'sampling' in message_lower: + return MockResponse("weighted") + elif 'card type' in message_lower: + return MockResponse("basic") + + # Default fallback + return MockResponse(default if default else (choices[0] if choices else "")) + with patch('rich.prompt.Prompt.ask', side_effect=mock_prompt_ask), \ patch('rich.prompt.IntPrompt.ask', side_effect=mock_int_prompt_ask), \ - patch('rich.prompt.Confirm.ask', side_effect=mock_confirm_ask): + patch('rich.prompt.Confirm.ask', side_effect=mock_confirm_ask), \ + patch('questionary.select', side_effect=mock_questionary_select): yield mock_responses @@ -146,7 +167,7 @@ def test_setup_when_config_missing(self, clean_temp_config, mock_services, mock_ def mock_prompt(text, **kwargs): if 'Obsidian' in text: return 'test_obs_key_123' - elif 'Anthropic' in text: + elif 'API key' in text or 'api key' in text.lower(): return 'test_anthro_key_456' elif 'Sampling mode' in text: return 'random' @@ -160,6 +181,23 @@ def mock_int_prompt(text, **kwargs): def mock_confirm(text, **kwargs): return kwargs.get('default', False) + def mock_questionary(message, choices=None, default=None, **kwargs): + """Mock questionary.select()""" + class MockResponse: + def __init__(self, value): + self.value = value + def ask(self): + return self.value + + message_lower = message.lower() + if 'model' in message_lower: + return MockResponse("Claude Sonnet 4.5") + elif 'sampling' in message_lower: + return MockResponse("random") + elif 'card type' in message_lower: + return MockResponse("basic") + return MockResponse(default if default else (choices[0] if choices else "")) + # Patch both the wizard module's paths and config module's paths with patch('obsidianki.cli.wizard.CONFIG_DIR', test_config_dir), \ patch('obsidianki.cli.wizard.ENV_FILE', test_env), \ @@ -168,7 +206,8 @@ def mock_confirm(text, **kwargs): patch('obsidianki.cli.config.CONFIG_FILE', test_config), \ patch('rich.prompt.Prompt.ask', side_effect=mock_prompt), \ patch('rich.prompt.IntPrompt.ask', side_effect=mock_int_prompt), \ - patch('rich.prompt.Confirm.ask', side_effect=mock_confirm): + patch('rich.prompt.Confirm.ask', side_effect=mock_confirm), \ + patch('questionary.select', side_effect=mock_questionary): from obsidianki.cli.wizard import setup setup(force_full_setup=True) @@ -195,7 +234,21 @@ def test_empty_api_keys_rejected(self, clean_temp_config, mock_services): def mock_empty_prompt(prompt_text, **kwargs): return "" # Empty key - with patch('rich.prompt.Prompt.ask', side_effect=mock_empty_prompt): + def mock_questionary_for_empty_test(message, choices=None, default=None, **kwargs): + """Mock questionary.select() for empty key test""" + class MockResponse: + def __init__(self, value): + self.value = value + def ask(self): + return self.value + + # Return valid values for model selection so test gets to API key prompt + if 'model' in message.lower(): + return MockResponse("Claude Sonnet 4.5") + return MockResponse(default if default else (choices[0] if choices else "")) + + with patch('rich.prompt.Prompt.ask', side_effect=mock_empty_prompt), \ + patch('questionary.select', side_effect=mock_questionary_for_empty_test): sys.argv = ['oki', '--setup'] from obsidianki.main import main