From f512a167d2d3e3c7fca5cbf6f7dbee40c26d3c0d Mon Sep 17 00:00:00 2001 From: FaustoS88 Date: Mon, 3 Aug 2026 23:25:21 +0200 Subject: [PATCH] fix(ci): pin ruff and declare explicit lint rule set MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The Lint job started failing on an unchanged commit: 5cf2aa1 passed on Jul 21 and failed on Aug 3 with 96 errors. The code never changed — ruff did. The workflow installed ruff unpinned and pyproject declared [tool.ruff] with no `select`, so the project inherited ruff's *default* rule set. ruff 0.16 shipped a much broader default (adding BLE001, ASYNC230, DTZ005, EXE001, I001, UP006, ...) and dropped E402, which also turned the existing `# noqa: E402` comments into RUF100 unused-noqa. Root-cause fix, so a future ruff release cannot break a green build: - pin ruff==0.16.1 in the Lint workflow - declare the rule set explicitly in pyproject: E, W, F, I, UP, B, C4, SIM, PIE, DTZ, RUF Resolved every finding under that set for real rather than suppressing: import sorting, PEP 585/604 annotations, datetime.UTC, tuple-form startswith, sorted() over a set, f-string conversion flags, combined with statements, and explicit zip(strict=True) where the lists are built in lockstep. Long lines wrapped instead of ignored. BLE001 and ASYNC230 are deliberately not selected: broad `except Exception` is intentional log-and-continue behaviour in the CLI scripts, and blocking file IO is fine in the one-shot crawler/eval scripts. Selecting them would have meant ~31 noqa comments across 12 files. Also completes the Pydantic AI v2 migration, which had left the suite at 16 failed / 89 passed: - Dependencies dropped openrouter_api_key/use_openrouter; tests updated and OpenRouter routing now covered via resolve_model() - agent introspection renamed: _function_tools -> the Capability's toolset, result_type -> output_type - retrieve() delegates to hybrid_retrieve, so its tests patch that boundary instead of mocking pool.fetch internals - chunk-header assertions updated v5 -> v6 - api_debug.py: OpenAIModel -> OpenAIChatModel, api_key moved onto OpenAIProvider, removed the no-longer-callable direct model invocation, result.data -> result.output Verified: ruff clean under both 0.16.1 and 0.15.4, 106 tests pass, all modules import. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01E9VRuDYsPBrRsbDCRrJVHb --- .github/workflows/lint.yml | 2 +- agent.py | 56 ++++---- api_debug.py | 72 +++++------ clear_database.py | 31 +++-- config.py | 12 +- db_inspect.py | 51 ++++---- db_schema.py | 22 ++-- init_db.py | 20 +-- interactive.py | 7 +- pinescript_crawler.py | 228 +++++++++++++++++---------------- pinescript_recrawl_light.py | 36 +++--- pyproject.toml | 15 +++ rag_utils.py | 19 +-- run.py | 54 +++++--- setup.py | 39 +++--- streamlit_ui.py | 23 ++-- tests/ragas_eval.py | 86 ++++++------- tests/test_agent.py | 185 +++++++++++++------------- tests/test_config.py | 16 ++- tests/test_models.py | 51 ++++---- tests/test_rag_improvements.py | 14 +- tests/test_rag_tier2.py | 24 ++-- 22 files changed, 562 insertions(+), 501 deletions(-) mode change 100644 => 100755 api_debug.py mode change 100644 => 100755 clear_database.py mode change 100644 => 100755 db_inspect.py mode change 100644 => 100755 init_db.py mode change 100644 => 100755 interactive.py mode change 100644 => 100755 run.py mode change 100644 => 100755 setup.py diff --git a/.github/workflows/lint.yml b/.github/workflows/lint.yml index e425b59..3392dc1 100644 --- a/.github/workflows/lint.yml +++ b/.github/workflows/lint.yml @@ -14,6 +14,6 @@ jobs: with: python-version: "3.11" - name: Install ruff - run: pip install ruff + run: pip install ruff==0.16.1 - name: Lint run: ruff check . diff --git a/agent.py b/agent.py index 1147347..6270c5c 100644 --- a/agent.py +++ b/agent.py @@ -16,32 +16,32 @@ import logging import os import sys +from collections.abc import AsyncGenerator from contextlib import asynccontextmanager from dataclasses import dataclass -from typing import AsyncGenerator import asyncpg +from dotenv import load_dotenv from openai import AsyncOpenAI from pydantic import BaseModel, Field from pydantic_ai import Agent, RunContext from pydantic_ai.capabilities import Capability, Thinking from pydantic_ai.models.openai import OpenAIChatModel from pydantic_ai.providers.openai import OpenAIProvider -from dotenv import load_dotenv from config import ( + DEFAULT_DATABASE_URL, DEFAULT_MODEL, EMBEDDING_MODEL, + HYBRID_SEARCH_ALPHA, LLM_MAX_TOKENS, LLM_TEMPERATURE, + MMR_LAMBDA, OPENROUTER_BASE_URL, OPENROUTER_DEFAULT_MODEL, - DEFAULT_DATABASE_URL, - HYBRID_SEARCH_ALPHA, - SIMILARITY_THRESHOLD, - RETRIEVAL_CANDIDATES, RERANK_TOP_N, - MMR_LAMBDA, + RETRIEVAL_CANDIDATES, + SIMILARITY_THRESHOLD, THINKING_EFFORT, ) from rag_utils import hybrid_retrieve @@ -59,15 +59,12 @@ # API key helper (unchanged from v1) # --------------------------------------------------------------------------- + def get_openai_api_key() -> str: """Get OpenAI API key with validation and interactive prompt if needed.""" openai_api_key = os.getenv("OPENAI_API_KEY") - if ( - not openai_api_key - or "YOUR_" in openai_api_key - or openai_api_key == "sk-..." - ): + if not openai_api_key or "YOUR_" in openai_api_key or openai_api_key == "sk-...": logger.warning("OPENAI_API_KEY is not set or has a placeholder value.") print("Please enter your OpenAI API key:") openai_api_key = input("> ") @@ -78,6 +75,7 @@ def get_openai_api_key() -> str: env_path = os.path.join(os.path.dirname(__file__), ".env") try: from dotenv import set_key + set_key(env_path, "OPENAI_API_KEY", openai_api_key) logger.info("Updated OPENAI_API_KEY in %s", env_path) except ImportError: @@ -90,6 +88,7 @@ def get_openai_api_key() -> str: # Output schema (unchanged) # --------------------------------------------------------------------------- + class PineScriptResult(BaseModel): query: str = Field(description="The original query") response: str = Field(description="The generated response") @@ -101,9 +100,11 @@ class PineScriptResult(BaseModel): # fields that were never read by any tool. # --------------------------------------------------------------------------- + @dataclass class Dependencies: """Runtime dependencies injected per run.""" + openai: AsyncOpenAI pool: asyncpg.Pool @@ -180,12 +181,11 @@ async def retrieve(ctx: RunContext[Dependencies], search_query: str) -> str: logger.debug("Hybrid retrieval returned %d documents", len(docs)) return "\n\n".join( - f"# {doc.title}\nDocumentation URL: {doc.url}\n\n{doc.content}\n" - for doc in docs + f"# {doc.title}\nDocumentation URL: {doc.url}\n\n{doc.content}\n" for doc in docs ) except Exception as e: logger.error("Error in retrieve tool: %s", e) - return f"Error retrieving documentation: {str(e)}" + return f"Error retrieving documentation: {e!s}" # --------------------------------------------------------------------------- @@ -199,16 +199,16 @@ async def retrieve(ctx: RunContext[Dependencies], search_query: str) -> str: # --------------------------------------------------------------------------- pinescript_agent = Agent( - DEFAULT_MODEL, # "openai:gpt-4o-mini" + DEFAULT_MODEL, # "openai:gpt-4o-mini" deps_type=Dependencies, - output_type=PineScriptResult, # v2 naming (was result_type in old v1) + output_type=PineScriptResult, # v2 naming (was result_type in old v1) capabilities=[ - Thinking(effort=THINKING_EFFORT), # extended reasoning, unified across providers - pinescript_rag, # retrieve tool + expert instructions + Thinking(effort=THINKING_EFFORT), # extended reasoning, unified across providers + pinescript_rag, # retrieve tool + expert instructions ], model_settings={ - "temperature": LLM_TEMPERATURE, # 0.2 — now actually applied every run - "max_tokens": LLM_MAX_TOKENS, # 2000 + "temperature": LLM_TEMPERATURE, # 0.2 — now actually applied every run + "max_tokens": LLM_MAX_TOKENS, # 2000 }, ) @@ -217,6 +217,7 @@ async def retrieve(ctx: RunContext[Dependencies], search_query: str) -> str: # Database helper (unchanged) # --------------------------------------------------------------------------- + @asynccontextmanager async def database_connect(create_db: bool = False) -> AsyncGenerator[asyncpg.Pool, None]: """Connect to the pgvector database.""" @@ -242,6 +243,7 @@ async def database_connect(create_db: bool = False) -> AsyncGenerator[asyncpg.Po # v2: resolve_model() → agent.run(model=...) (no context manager needed) # --------------------------------------------------------------------------- + def resolve_model(preset: str | None = None) -> OpenAIChatModel | None: """Resolve a preset name or raw model ID to a Pydantic AI model. @@ -261,9 +263,7 @@ def resolve_model(preset: str | None = None) -> OpenAIChatModel | None: logger.info("Using raw model ID: %s", model_id) if openrouter_api_key: - provider = OpenAIProvider( - base_url=OPENROUTER_BASE_URL, api_key=openrouter_api_key - ) + provider = OpenAIProvider(base_url=OPENROUTER_BASE_URL, api_key=openrouter_api_key) return OpenAIChatModel(model_id, provider=provider) logger.warning("No OPENROUTER_API_KEY for preset '%s', using default", preset) @@ -271,9 +271,7 @@ def resolve_model(preset: str | None = None) -> OpenAIChatModel | None: # No preset: use OpenRouter default if key exists, else agent default if openrouter_api_key: - provider = OpenAIProvider( - base_url=OPENROUTER_BASE_URL, api_key=openrouter_api_key - ) + provider = OpenAIProvider(base_url=OPENROUTER_BASE_URL, api_key=openrouter_api_key) return OpenAIChatModel(OPENROUTER_DEFAULT_MODEL, provider=provider) return None # agent default (OpenAI) @@ -283,6 +281,7 @@ def resolve_model(preset: str | None = None) -> OpenAIChatModel | None: # Run helper — now accepts message_history for proper multi-turn # --------------------------------------------------------------------------- + async def run_agent(question: str, preset: str | None = None, message_history=None): """Run the agent with a specific question. @@ -322,6 +321,7 @@ async def run_agent(question: str, preset: str | None = None, message_history=No # CLI entry point # --------------------------------------------------------------------------- + async def main(): """Run the agent from the command line.""" if len(sys.argv) > 1: @@ -334,7 +334,7 @@ async def main(): if result: print("\nResponse:") - print(result.output.response) # v2: .output (was .data) + print(result.output.response) # v2: .output (was .data) print(f"\nSnippets used: {result.output.snippets_used}") else: print("No response received from the agent.") diff --git a/api_debug.py b/api_debug.py old mode 100644 new mode 100755 index 6f909fd..3191b76 --- a/api_debug.py +++ b/api_debug.py @@ -6,42 +6,46 @@ different aspects of the initialization and request process. """ -import os import asyncio +import os + from dotenv import load_dotenv -from openai import OpenAI, AsyncOpenAI -from pydantic_ai.models.openai import OpenAIModel +from openai import AsyncOpenAI, OpenAI from pydantic_ai import Agent +from pydantic_ai.models.openai import OpenAIChatModel +from pydantic_ai.providers.openai import OpenAIProvider # Load environment variables load_dotenv(override=True) + def print_section(title): """Print a section title""" print("\n" + "=" * 80) print(f" {title} ".center(80, "=")) print("=" * 80) + async def main(): """Main function to debug API key issues""" print_section("Environment Variables") - + # Check if API key is set in environment api_key = os.getenv("OPENAI_API_KEY") masked_key = f"{api_key[:4]}...{api_key[-4:]}" if api_key and len(api_key) > 8 else "None" print(f"OPENAI_API_KEY from environment: {masked_key}") - + # Check for malformed placeholders if api_key in ["YOUR_OPENAI_API_KEY", "sk-...", "YOUR_OPE***_API"] or not api_key: print("WARNING: API key appears to be a placeholder or is missing!") - + # Prompt for key print("Enter your OpenAI API key for testing:") api_key = input("> ").strip() os.environ["OPENAI_API_KEY"] = api_key masked_key = f"{api_key[:4]}...{api_key[-4:]}" if api_key and len(api_key) > 8 else "None" print(f"Using API key: {masked_key}") - + # Test standard OpenAI client print_section("Standard OpenAI Client Test") try: @@ -51,7 +55,7 @@ async def main(): print(f"Found {len(models.data)} models") except Exception as e: print(f"❌ Standard OpenAI client failed: {e}") - + # Test AsyncOpenAI client print_section("Async OpenAI Client Test") try: @@ -61,50 +65,33 @@ async def main(): print(f"Found {len(models.data)} models") except Exception as e: print(f"❌ Async OpenAI client failed: {e}") - - # Test Pydantic AI OpenAIModel - print_section("Pydantic AI OpenAIModel Test") + + # Test Pydantic AI OpenAIChatModel + # v2: the API key travels on the provider, not as a model kwarg. + print_section("Pydantic AI OpenAIChatModel Test") + openai_model = None try: - openai_model = OpenAIModel("gpt-4o", api_key=api_key) - print("✅ OpenAIModel initialized successfully!") + openai_model = OpenAIChatModel("gpt-4o", provider=OpenAIProvider(api_key=api_key)) + print("✅ OpenAIChatModel initialized successfully!") print(f"Model name: {openai_model.model_name}") - - # Test making a request with the model - response = await openai_model("Hello, world!") - print("✅ Model request successful!") - print(f"Response: {response[:50]}...") except Exception as e: - print(f"❌ OpenAIModel failed: {e}") - + print(f"❌ OpenAIChatModel failed: {e}") + # Test Pydantic AI Agent print_section("Pydantic AI Agent Test") try: - # Try creating a minimal agent - agent = Agent( - model="openai:gpt-4o", - model_settings={ - "api_key": api_key - } - ) + # v2: pass the configured model object; api_key is not a model_setting. + agent = Agent(openai_model) if openai_model else Agent("openai:gpt-4o") print("✅ Agent created successfully!") - - # Print the agent's model settings - print("Agent model settings:") - if hasattr(agent, 'model_settings'): - for key, value in agent.model_settings.items(): - if key == "api_key" and value: - print(f" api_key: {value[:4]}...{value[-4:]}") - else: - print(f" {key}: {value}") - + # Try making a request with the agent try: - result = agent.run_sync("Hello, world!", model_settings={"api_key": api_key}) + result = agent.run_sync("Hello, world!") print("✅ Agent request successful!") - print(f"Response: {result.data[:50]}...") + print(f"Response: {str(result.output)[:50]}...") except Exception as e: print(f"❌ Agent request failed: {e}") - + # Check if the error is related to API key error_str = str(e).lower() if "api key" in error_str or "openai" in error_str: @@ -115,7 +102,7 @@ async def main(): print("3. Modify the retrieve tool to use a separate OpenAI client") except Exception as e: print(f"❌ Agent creation failed: {e}") - + print_section("Summary") print("This debug information should help identify where the API key issue is occurring.") print("Look for any failures above and focus debugging on those components.") @@ -124,5 +111,6 @@ async def main(): print("2. If some tests fail, focus on those specific components") print("3. Check if `model_settings` in the agent includes the correct API key") + if __name__ == "__main__": - asyncio.run(main()) \ No newline at end of file + asyncio.run(main()) diff --git a/clear_database.py b/clear_database.py old mode 100644 new mode 100755 index 5800811..5a69a80 --- a/clear_database.py +++ b/clear_database.py @@ -7,17 +7,20 @@ """ import asyncio + from dotenv import load_dotenv + from agent import database_connect from db_schema import validate_schema # Load environment variables load_dotenv() + async def clear_database(): """Clear all records from the pinescript_docs table""" print("\n===== DATABASE CLEANUP UTILITY =====\n") - + # Verify database connection and structure async with database_connect(False) as pool: # Validate schema first @@ -25,47 +28,49 @@ async def clear_database(): if not schema_valid: print("Database schema is not valid, please run init_db.py first.") return False - + # Count total documents count = await pool.fetchval("SELECT COUNT(*) FROM pinescript_docs") - + if count == 0: print("Database is already empty.") return True - + print(f"Found {count} records in the database.") print("This operation will delete ALL records from the pinescript_docs table.") - + # Get user confirmation confirmation = input("Type 'DELETE' to confirm deletion: ") - - if confirmation.upper() != 'DELETE': + + if confirmation.upper() != "DELETE": print("Operation cancelled.") return False - + # Execute deletion try: await pool.execute("DELETE FROM pinescript_docs") - + # Verify deletion new_count = await pool.fetchval("SELECT COUNT(*) FROM pinescript_docs") - + print(f"Successfully deleted {count} records.") print(f"Database now contains {new_count} records.") - + return True except Exception as e: print(f"Error deleting records: {e}") return False + async def main(): """Main function""" success = await clear_database() - + if success: print("\nDatabase cleared successfully. Ready for a fresh crawl.") else: print("\nFailed to clear database. Please check the error messages.") + if __name__ == "__main__": - asyncio.run(main()) \ No newline at end of file + asyncio.run(main()) diff --git a/config.py b/config.py index 5c80748..9a98248 100644 --- a/config.py +++ b/config.py @@ -41,9 +41,7 @@ RERANK_TOP_N = int(os.getenv("PINESCRIPT_RERANK_TOP_N", "15")) # Cross-encoder model for reranking -RERANK_MODEL = os.getenv( - "PINESCRIPT_RERANK_MODEL", "cross-encoder/ms-marco-MiniLM-L-6-v2" -) +RERANK_MODEL = os.getenv("PINESCRIPT_RERANK_MODEL", "cross-encoder/ms-marco-MiniLM-L-6-v2") # MMR diversity parameter (1.0 = pure relevance, 0.0 = pure diversity) MMR_LAMBDA = float(os.getenv("PINESCRIPT_MMR_LAMBDA", "0.7")) @@ -71,12 +69,8 @@ # OpenRouter settings # --------------------------------------------------------------------------- -OPENROUTER_BASE_URL = os.getenv( - "OPENROUTER_BASE_URL", "https://openrouter.ai/api/v1" -) -OPENROUTER_DEFAULT_MODEL = os.getenv( - "OPENROUTER_MODEL", "openai/gpt-4.1-mini" -) +OPENROUTER_BASE_URL = os.getenv("OPENROUTER_BASE_URL", "https://openrouter.ai/api/v1") +OPENROUTER_DEFAULT_MODEL = os.getenv("OPENROUTER_MODEL", "openai/gpt-4.1-mini") # --------------------------------------------------------------------------- # Database diff --git a/db_inspect.py b/db_inspect.py old mode 100644 new mode 100755 index ba878d4..39ca0c9 --- a/db_inspect.py +++ b/db_inspect.py @@ -6,56 +6,60 @@ check vector quality, and perform test searches without running the full agent. """ +import asyncio import os import sys -import asyncio + import pydantic_core from dotenv import load_dotenv -from agent import database_connect from openai import AsyncOpenAI +from agent import database_connect + # Load environment variables load_dotenv() + async def count_entries(): """Count the number of entries in the database""" async with database_connect(False) as pool: count = await pool.fetchval("SELECT COUNT(*) FROM pinescript_docs") print(f"Database contains {count} documentation sections") + async def list_titles(limit=20): """List document titles in the database""" async with database_connect(False) as pool: rows = await pool.fetch( - "SELECT id, title, url FROM pinescript_docs ORDER BY id LIMIT $1", - limit + "SELECT id, title, url FROM pinescript_docs ORDER BY id LIMIT $1", limit ) - + print(f"Sample of {len(rows)} document titles:") for row in rows: print(f"[{row['id']}] {row['title']}") print(f" URL: {row['url']}") print() + async def view_document(doc_id): """View a specific document by ID""" async with database_connect(False) as pool: row = await pool.fetchrow( - "SELECT id, title, url, content FROM pinescript_docs WHERE id = $1", - doc_id + "SELECT id, title, url, content FROM pinescript_docs WHERE id = $1", doc_id ) - + if not row: print(f"No document found with ID {doc_id}") return - + print(f"Document #{row['id']}: {row['title']}") print(f"URL: {row['url']}") print("\nContent:") print("=" * 80) - print(row['content']) + print(row["content"]) print("=" * 80) + async def test_search(query): """Test search functionality""" # Initialize OpenAI client @@ -63,9 +67,9 @@ async def test_search(query): if not openai_api_key: print("Error: OPENAI_API_KEY not found in environment variables") return - + openai = AsyncOpenAI(api_key=openai_api_key) - + # Generate embedding for query print(f"Generating embedding for query: '{query}'") embedding = await openai.embeddings.create( @@ -74,7 +78,7 @@ async def test_search(query): ) embedding_vector = embedding.data[0].embedding embedding_json = pydantic_core.to_json(embedding_vector).decode() - + # Search database async with database_connect(False) as pool: print("Searching database...") @@ -88,46 +92,48 @@ async def test_search(query): """, embedding_json, ) - + print(f"\nTop {len(rows)} results for '{query}':") print("=" * 80) for i, row in enumerate(rows): - print(f"Result #{i+1} [Similarity: {row['similarity']:.4f}]") + print(f"Result #{i + 1} [Similarity: {row['similarity']:.4f}]") print(f"Title: {row['title']}") print(f"URL: {row['url']}") print(f"Content Preview: {row['content'][:150]}...") print("-" * 80) + async def verify_vectors(): """Verify that all documents have valid embedding vectors""" async with database_connect(False) as pool: # Count total documents total = await pool.fetchval("SELECT COUNT(*) FROM pinescript_docs") - + # Check for NULL embeddings null_embeddings = await pool.fetchval( "SELECT COUNT(*) FROM pinescript_docs WHERE embedding IS NULL" ) - + # Check for zero-length embeddings zero_length = await pool.fetchval( "SELECT COUNT(*) FROM pinescript_docs WHERE array_length(embedding, 1) = 0" ) - + # Get dimension of embeddings dimension = await pool.fetchval( "SELECT array_length(embedding, 1) FROM pinescript_docs LIMIT 1" ) - + print(f"Total documents: {total}") print(f"Documents with NULL embeddings: {null_embeddings}") print(f"Documents with zero-length embeddings: {zero_length}") print(f"Embedding dimension: {dimension}") - + if null_embeddings > 0 or zero_length > 0: print("\nWarning: Some documents have invalid embeddings!") print("This may affect search quality. Consider regenerating embeddings.") + async def main(): """Main function""" if len(sys.argv) < 2: @@ -138,9 +144,9 @@ async def main(): print(" python db_inspect.py search - Test search functionality") print(" python db_inspect.py verify - Verify vector quality") return - + command = sys.argv[1] - + if command == "count": await count_entries() elif command == "list": @@ -161,5 +167,6 @@ async def main(): else: print(f"Unknown command: {command}") + if __name__ == "__main__": asyncio.run(main()) diff --git a/db_schema.py b/db_schema.py index b509fe0..5dbcf80 100644 --- a/db_schema.py +++ b/db_schema.py @@ -61,13 +61,14 @@ ALTER TABLE pinescript_docs ADD COLUMN IF NOT EXISTS contextual_prefix TEXT; """ + # Function to validate the schema (can be called to ensure DB is ready) async def validate_schema(pool): """Validate that the database schema is correctly set up. - + Args: pool: Database connection pool - + Returns: bool: Whether the schema is valid """ @@ -79,7 +80,7 @@ async def validate_schema(pool): if not extension_exists: print("Error: pgvector extension is not installed") return False - + # Check if the pinescript_docs table exists table_exists = await pool.fetchval( "SELECT 1 FROM information_schema.tables WHERE table_name = 'pinescript_docs'" @@ -87,7 +88,7 @@ async def validate_schema(pool): if not table_exists: print("Error: pinescript_docs table does not exist") return False - + # Check if the embedding index exists index_exists = await pool.fetchval( "SELECT 1 FROM pg_indexes WHERE indexname = 'idx_pinescript_docs_embedding'" @@ -95,12 +96,13 @@ async def validate_schema(pool): if not index_exists: print("Warning: idx_pinescript_docs_embedding index does not exist") # Not returning False here as missing index is not critical - + return True except Exception as e: print(f"Error validating schema: {e}") return False + # Function to create the schema async def create_schema(pool): """Create the database schema. @@ -112,9 +114,8 @@ async def create_schema(pool): bool: Whether the schema was created successfully """ try: - async with pool.acquire() as conn: - async with conn.transaction(): - await conn.execute(DB_SCHEMA) + async with pool.acquire() as conn, conn.transaction(): + await conn.execute(DB_SCHEMA) print("Database schema created successfully") return True except Exception as e: @@ -134,9 +135,8 @@ async def run_migration(pool): bool: Whether the migration ran successfully """ try: - async with pool.acquire() as conn: - async with conn.transaction(): - await conn.execute(DB_MIGRATION) + async with pool.acquire() as conn, conn.transaction(): + await conn.execute(DB_MIGRATION) print("Migration completed successfully") return True except Exception as e: diff --git a/init_db.py b/init_db.py old mode 100644 new mode 100755 index cc582fe..1bef86a --- a/init_db.py +++ b/init_db.py @@ -10,9 +10,10 @@ Run this script before running the crawler or using the agent. """ +import asyncio import logging import os -import asyncio + import asyncpg from dotenv import load_dotenv @@ -23,6 +24,7 @@ logger = logging.getLogger(__name__) + async def init_database(): """Initialize the database""" # Get database connection parameters @@ -46,20 +48,16 @@ async def init_database(): ) logger.info("Tables in database:") for table in tables: - logger.info(" - %s", table['table_name']) + logger.info(" - %s", table["table_name"]) # Verify vector extension - extensions = await conn.fetch( - "SELECT extname FROM pg_extension" - ) + extensions = await conn.fetch("SELECT extname FROM pg_extension") logger.info("Extensions in database:") for ext in extensions: - logger.info(" - %s", ext['extname']) + logger.info(" - %s", ext["extname"]) # Check if vector extension is installed - vector_ext = await conn.fetchval( - "SELECT 1 FROM pg_extension WHERE extname = 'vector'" - ) + vector_ext = await conn.fetchval("SELECT 1 FROM pg_extension WHERE extname = 'vector'") if vector_ext: logger.info("Vector extension is installed") else: @@ -71,14 +69,16 @@ async def init_database(): except Exception as e: logger.error("Error initializing database: %s", e) finally: - if 'conn' in locals(): + if "conn" in locals(): await conn.close() logger.info("Database connection closed") + def main(): """Main entry point""" logging.basicConfig(level=logging.INFO, format="%(levelname)s %(message)s") asyncio.run(init_database()) + if __name__ == "__main__": main() diff --git a/interactive.py b/interactive.py old mode 100644 new mode 100755 index a29599d..f90fdb3 --- a/interactive.py +++ b/interactive.py @@ -8,9 +8,10 @@ import asyncio import sys + from dotenv import load_dotenv -from agent import run_agent, PineScriptResult, get_openai_api_key +from agent import PineScriptResult, get_openai_api_key, run_agent load_dotenv(override=True) @@ -66,7 +67,7 @@ async def process_input(self, user_input: str) -> bool: if result and isinstance(result.output, PineScriptResult): print("\n" + "=" * 80) - print(result.output.response) # v2: .output (was .data) + print(result.output.response) # v2: .output (was .data) print("=" * 80) # Accumulate history so the next turn sees all prior context @@ -102,4 +103,4 @@ async def main(preset: str | None = None): if __name__ == "__main__": - asyncio.run(main()) \ No newline at end of file + asyncio.run(main()) diff --git a/pinescript_crawler.py b/pinescript_crawler.py index 52fd3d9..fe44745 100644 --- a/pinescript_crawler.py +++ b/pinescript_crawler.py @@ -3,29 +3,28 @@ import asyncio import logging import os -import sys import re -import pydantic_core -from typing import List, Dict, Set +import sys +from datetime import UTC, datetime import asyncpg -from openai import AsyncOpenAI -from dotenv import load_dotenv +import pydantic_core from bs4 import BeautifulSoup -from datetime import datetime # Import crawler from crawl4ai from crawl4ai import AsyncWebCrawler, BrowserConfig from crawl4ai.extraction_strategy import JsonCssExtractionStrategy +from dotenv import load_dotenv +from openai import AsyncOpenAI # Force reload environment variables load_dotenv(override=True) # Import database setup functions -from db_schema import create_schema, run_migration # noqa: E402 from agent import database_connect # noqa: E402 +from config import CHUNK_OVERLAP, CHUNK_SIZE # noqa: E402 +from db_schema import create_schema, run_migration # noqa: E402 from rag_utils import prepend_chunk_header, recursive_character_split # noqa: E402 -from config import CHUNK_SIZE, CHUNK_OVERLAP # noqa: E402 logging.basicConfig( level=logging.INFO, @@ -42,7 +41,7 @@ def __init__(self): # This is the correct base URL that works in the original script self.base_url = "https://www.tradingview.com/pine-script-docs" self.output_dir = "pinescript_docs" - self.visited_urls: Set[str] = set() + self.visited_urls: set[str] = set() # Initialize OpenAI client for embeddings openai_api_key = os.getenv("OPENAI_API_KEY") @@ -66,27 +65,23 @@ def __init__(self): "name": "PineScript Documentation", "baseSelector": "main", # Main content area "fields": [ - { - "name": "title", - "selector": "h1", - "type": "text" - }, + {"name": "title", "selector": "h1", "type": "text"}, { "name": "content", "selector": "main > div", # Main content excluding navigation - "type": "html" + "type": "html", }, { "name": "navigation", "selector": "nav", # Left navigation menu - "type": "html" + "type": "html", }, { "name": "toc", "selector": "[aria-label='Table of contents']", # Right-side TOC - "type": "html" - } - ] + "type": "html", + }, + ], } # Semaphore to limit concurrent embedding API calls @@ -98,32 +93,31 @@ def normalize_url(self, url: str) -> str: return "" # Remove anchor tags and query parameters - url = url.split('#')[0].split('?')[0] + url = url.split("#")[0].split("?")[0] # Skip external links and special protocols - if url.startswith(('http', 'https')) and not url.startswith(self.base_url): + if url.startswith(("http", "https")) and not url.startswith(self.base_url): return "" - if url.startswith(('mailto:', 'tel:', 'javascript:')): + if url.startswith(("mailto:", "tel:", "javascript:")): return "" # Handle relative URLs - if not url.startswith('http'): - if url.startswith('/'): + if not url.startswith("http"): + if url.startswith("/"): url = f"https://www.tradingview.com{url}" else: url = f"{self.base_url}/{url}" return url - async def get_all_doc_urls(self) -> List[str]: + async def get_all_doc_urls(self) -> list[str]: """Extract all documentation URLs from the navigation menu""" urls = set() logger.info("Starting to collect URLs...") # Start with main sections from left navigation browser_config = BrowserConfig( - headless=True, - extra_args=["--disable-gpu", "--disable-dev-shm-usage", "--no-sandbox"] + headless=True, extra_args=["--disable-gpu", "--disable-dev-shm-usage", "--no-sandbox"] ) # THIS IS CRITICAL - the /welcome/ path is the entry point that works @@ -134,13 +128,13 @@ async def get_all_doc_urls(self) -> List[str]: result = await crawler.arun(url=welcome_url) if result.success: logger.info("Successfully accessed the main page") - soup = BeautifulSoup(result.html, 'html.parser') + soup = BeautifulSoup(result.html, "html.parser") # Find all navigation elements - nav_elements = soup.find_all(['nav', 'div'], class_=['toc', 'sidebar']) + nav_elements = soup.find_all(["nav", "div"], class_=["toc", "sidebar"]) for nav in nav_elements: - for link in nav.find_all('a'): - href = link.get('href') + for link in nav.find_all("a"): + href = link.get("href") if href: full_url = self.normalize_url(href) if full_url: @@ -156,13 +150,13 @@ async def get_all_doc_urls(self) -> List[str]: fallback_result = await crawler.arun(url=fallback_url) if fallback_result.success: logger.info("Successfully accessed fallback page") - soup = BeautifulSoup(fallback_result.html, 'html.parser') + soup = BeautifulSoup(fallback_result.html, "html.parser") # Find all navigation elements - nav_elements = soup.find_all(['nav', 'div'], class_=['toc', 'sidebar']) + nav_elements = soup.find_all(["nav", "div"], class_=["toc", "sidebar"]) for nav in nav_elements: - for link in nav.find_all('a'): - href = link.get('href') + for link in nav.find_all("a"): + href = link.get("href") if href: full_url = self.normalize_url(href) if full_url: @@ -173,14 +167,20 @@ async def get_all_doc_urls(self) -> List[str]: if not urls: logger.warning("No URLs found in navigation, trying common paths...") common_paths = [ - "welcome/", "introduction/", "concepts/", "language/", - "essential/", "resources/", "reference/", "faq/" + "welcome/", + "introduction/", + "concepts/", + "language/", + "essential/", + "resources/", + "reference/", + "faq/", ] for path in common_paths: urls.add(f"{self.base_url}/{path}") - urls_list = sorted(list(urls)) + urls_list = sorted(urls) logger.info("Total URLs found: %d", len(urls_list)) return urls_list @@ -189,20 +189,20 @@ async def get_all_doc_urls(self) -> List[str]: def clean_navigation(self, text): """Remove navigation elements and links""" # Remove navigation sections - text = re.sub(r'Version Version.*?Auto', '', text, flags=re.DOTALL) - text = re.sub(r'\* \[.*?\n', '', text) - text = re.sub(r'Copyright © .*?TradingView.*?\n', '', text) - text = re.sub(r'On this page.*?\n', '', text) + text = re.sub(r"Version Version.*?Auto", "", text, flags=re.DOTALL) + text = re.sub(r"\* \[.*?\n", "", text) + text = re.sub(r"Copyright © .*?TradingView.*?\n", "", text) + text = re.sub(r"On this page.*?\n", "", text) # Clean up additional navigation elements - text = re.sub(r'\[ User Manual \].*?\n', '', text) - text = re.sub(r'\[ Previous .*? \]', '', text) - text = re.sub(r'\[ Next .*? \]', '', text) + text = re.sub(r"\[ User Manual \].*?\n", "", text) + text = re.sub(r"\[ Previous .*? \]", "", text) + text = re.sub(r"\[ Next .*? \]", "", text) return text def extract_code_blocks(self, text): """Preserve and clean code blocks""" # Find Pine Script code blocks - code_blocks = re.findall(r'```(?:pine)?(.*?)```', text, re.DOTALL) + code_blocks = re.findall(r"```(?:pine)?(.*?)```", text, re.DOTALL) clean_blocks = [] for block in code_blocks: # Clean the code block @@ -214,13 +214,13 @@ def extract_code_blocks(self, text): def extract_function_docs(self, text): """Extract function documentation""" # Find function descriptions - functions = re.findall(r'@function.*?@returns.*?\n', text, re.DOTALL) + functions = re.findall(r"@function.*?@returns.*?\n", text, re.DOTALL) return functions def process_content(self, raw_content, url): """Process raw content into a cleaner format for embedding""" # Skip if no real content - if len(raw_content) < 100 or 'User Manual' not in raw_content: + if len(raw_content) < 100 or "User Manual" not in raw_content: return raw_content # Clean navigation and basic structure @@ -231,16 +231,19 @@ def process_content(self, raw_content, url): function_docs = self.extract_function_docs(content) # Extract main content sections (Q&A format in FAQ) - sections = re.findall(r'##\s+\[(.*?)\].*?\n(.*?)(?=##|\Z)', content, re.DOTALL) + sections = re.findall(r"##\s+\[(.*?)\].*?\n(.*?)(?=##|\Z)", content, re.DOTALL) # Build processed content processed = [] if sections: for title, section in sections: - if any(keyword in section.lower() for keyword in ['pine', 'script', 'function', 'indicator', 'value', 'parameter']): - clean_section = re.sub(r'\[\^.*?\]', '', section) # Remove footnotes - clean_section = re.sub(r'\(https://.*?\)', '', clean_section) # Remove links + if any( + keyword in section.lower() + for keyword in ["pine", "script", "function", "indicator", "value", "parameter"] + ): + clean_section = re.sub(r"\[\^.*?\]", "", section) # Remove footnotes + clean_section = re.sub(r"\(https://.*?\)", "", clean_section) # Remove links processed.append(f"## {title}\n{clean_section.strip()}") if code_blocks: @@ -257,7 +260,7 @@ def process_content(self, raw_content, url): return "\n\n".join(processed) - async def generate_embedding(self, text: str) -> List[float]: + async def generate_embedding(self, text: str) -> list[float]: """Generate embedding for text using OpenAI API""" async with self.sem: try: @@ -268,30 +271,25 @@ async def generate_embedding(self, text: str) -> List[float]: text = text[:max_length] response = await self.openai.embeddings.create( - input=text, - model="text-embedding-3-small" + input=text, model="text-embedding-3-small" ) return response.data[0].embedding except Exception as e: logger.error("Error generating embedding: %s", e) return [] - async def crawl_docs(self, urls: List[str]): + async def crawl_docs(self, urls: list[str]): """Crawl documentation pages and store in vector database""" logger.info("Starting crawling process...") # Configure extraction strategy - structure_strategy = JsonCssExtractionStrategy( - schema=self.structure_schema, - verbose=True - ) + structure_strategy = JsonCssExtractionStrategy(schema=self.structure_schema, verbose=True) browser_config = BrowserConfig( - headless=True, - extra_args=["--disable-gpu", "--disable-dev-shm-usage", "--no-sandbox"] + headless=True, extra_args=["--disable-gpu", "--disable-dev-shm-usage", "--no-sandbox"] ) - timestamp = datetime.now().strftime("%Y%m%d_%H%M%S") + timestamp = datetime.now(UTC).strftime("%Y%m%d_%H%M%S") combined_path = f"{self.output_dir}/all_docs_{timestamp}.md" failed_path = f"{self.output_dir}/failed_urls_{timestamp}.txt" @@ -306,13 +304,14 @@ async def crawl_docs(self, urls: List[str]): success = 0 failed = 0 - with open(combined_path, "w", encoding="utf-8") as combined_file, \ - open(failed_path, "w", encoding="utf-8") as failed_file: - + with ( + open(combined_path, "w", encoding="utf-8") as combined_file, + open(failed_path, "w", encoding="utf-8") as failed_file, + ): # Process in small batches batch_size = 3 for i in range(0, len(urls), batch_size): - batch = urls[i:i + batch_size] + batch = urls[i : i + batch_size] total_batches = (len(urls) + batch_size - 1) // batch_size logger.info("Processing batch %d/%d", i // batch_size + 1, total_batches) @@ -327,20 +326,21 @@ async def crawl_docs(self, urls: List[str]): # Use the extraction strategy result = await crawler.arun( - url=url, - extraction_strategy=structure_strategy + url=url, extraction_strategy=structure_strategy ) if result.success: # Get page name for file - page_name = url.rstrip('/').split('/')[-1] or 'index' + page_name = url.rstrip("/").split("/")[-1] or "index" file_path = f"{self.output_dir}/{page_name}_{timestamp}.md" - # Get the markdown content - handle both string and MarkdownGenerationResult object + # Handle both str and MarkdownGenerationResult if isinstance(result.markdown, str): markdown = result.markdown else: - markdown = result.markdown.raw_markdown if result.markdown else "" + markdown = ( + result.markdown.raw_markdown if result.markdown else "" + ) # Save individual file with open(file_path, "w", encoding="utf-8") as f: @@ -360,13 +360,15 @@ async def crawl_docs(self, urls: List[str]): success += 1 logger.info("Successfully processed: %s", page_name) else: - logger.warning("Failed to crawl %s: %s", url, result.error_message) + logger.warning( + "Failed to crawl %s: %s", url, result.error_message + ) failed_file.write(f"{url}: {result.error_message}\n") failed += 1 except Exception as e: logger.error("Error processing %s: %s", url, e) - failed_file.write(f"{url}: {str(e)}\n") + failed_file.write(f"{url}: {e!s}\n") failed += 1 # Rate limiting between batches @@ -376,7 +378,9 @@ async def crawl_docs(self, urls: List[str]): count = await pool.fetchval("SELECT COUNT(*) FROM pinescript_docs") logger.info( "Crawling completed — processed: %d, failed: %d, db sections: %d", - success, failed, count, + success, + failed, + count, ) async def process_and_store_document(self, url: str, markdown: str, pool: asyncpg.Pool): @@ -401,12 +405,11 @@ async def process_and_store_document(self, url: str, markdown: str, pool: asyncp for section in sections: # Check if already exists exists = await pool.fetchval( - "SELECT 1 FROM pinescript_docs WHERE url = $1", - section["url"] + "SELECT 1 FROM pinescript_docs WHERE url = $1", section["url"] ) if exists: - logger.debug("Section already exists: %s", section['url']) + logger.debug("Section already exists: %s", section["url"]) continue # Generate embedding @@ -414,7 +417,7 @@ async def process_and_store_document(self, url: str, markdown: str, pool: asyncp embedding = await self.generate_embedding(embedding_text) if not embedding: - logger.warning("Failed to generate embedding for %s", section['url']) + logger.warning("Failed to generate embedding for %s", section["url"]) continue # Convert to JSON @@ -433,9 +436,9 @@ async def process_and_store_document(self, url: str, markdown: str, pool: asyncp embedding_json, ) - logger.debug("Inserted section: %s", section['title']) + logger.debug("Inserted section: %s", section["title"]) - def split_into_sections(self, markdown: str, url: str) -> List[Dict[str, str]]: + def split_into_sections(self, markdown: str, url: str) -> list[dict[str, str]]: """Split markdown into chunks using recursive character splitting with overlap. Uses heading structure to identify sections, then applies recursive @@ -452,20 +455,22 @@ def split_into_sections(self, markdown: str, url: str) -> List[Dict[str, str]]: lines = lines[1:] # First pass: identify heading-based sections - raw_sections: List[Dict[str, str]] = [] + raw_sections: list[dict[str, str]] = [] current_section = None - current_content: List[str] = [] + current_content: list[str] = [] for line in lines: - if line.startswith("## ") or line.startswith("### "): + if line.startswith(("## ", "### ")): # Save previous section if current_section and current_content: content = "\n".join(current_content).strip() if content: - raw_sections.append({ - "title": current_section, - "content": content, - }) + raw_sections.append( + { + "title": current_section, + "content": content, + } + ) # Start new section heading = line.lstrip("#").strip() current_section = heading @@ -480,28 +485,34 @@ def split_into_sections(self, markdown: str, url: str) -> List[Dict[str, str]]: if current_section and current_content: content = "\n".join(current_content).strip() if content: - raw_sections.append({ - "title": current_section, - "content": content, - }) + raw_sections.append( + { + "title": current_section, + "content": content, + } + ) # Handle pre-heading content if not current_section and current_content: content = "\n".join(current_content).strip() if content: - raw_sections.append({ - "title": page_title, - "content": content, - }) + raw_sections.append( + { + "title": page_title, + "content": content, + } + ) # If no sections found at all, treat entire doc as one section if not raw_sections and lines: content = "\n".join(lines).strip() if content: - raw_sections.append({ - "title": page_title, - "content": content, - }) + raw_sections.append( + { + "title": page_title, + "content": content, + } + ) # Second pass: recursive character split within each section chunk_idx = 0 @@ -520,12 +531,14 @@ def split_into_sections(self, markdown: str, url: str) -> List[Dict[str, str]]: # Prepend contextual chunk header enriched = prepend_chunk_header(page_title, section_title, chunk) - section_id = re.sub(r'[^a-z0-9]+', '-', section_title.lower()) - sections.append({ - "url": f"{url}#{section_id}-{chunk_idx}", - "title": f"{page_title} - {section_title}", - "content": enriched, - }) + section_id = re.sub(r"[^a-z0-9]+", "-", section_title.lower()) + sections.append( + { + "url": f"{url}#{section_id}-{chunk_idx}", + "title": f"{page_title} - {section_title}", + "content": enriched, + } + ) chunk_idx += 1 return sections @@ -558,7 +571,8 @@ async def clear_database(): count_after = await pool.fetchval("SELECT COUNT(*) FROM pinescript_docs") logger.info( "Deleted %d records — database now contains %d", - count_before - count_after, count_after, + count_before - count_after, + count_after, ) diff --git a/pinescript_recrawl_light.py b/pinescript_recrawl_light.py index 41e2887..8d17eda 100644 --- a/pinescript_recrawl_light.py +++ b/pinescript_recrawl_light.py @@ -19,6 +19,7 @@ import os import re import sys + import asyncpg import html2text import httpx @@ -30,22 +31,22 @@ load_dotenv(override=True) # Import from existing project modules -from db_schema import create_schema, run_migration # noqa: E402 from agent import database_connect # noqa: E402 -from rag_utils import ( # noqa: E402 - code_aware_split, - detect_content_type, - generate_contextual_prefix, - prepend_chunk_header, -) from config import ( # noqa: E402 - CHUNK_SIZE, CHUNK_OVERLAP, + CHUNK_SIZE, CONTEXTUAL_MODEL, EMBEDDING_MODEL, MAX_PAGE_CONTEXT_CHARS, OPENROUTER_BASE_URL, ) +from db_schema import create_schema, run_migration # noqa: E402 +from rag_utils import ( # noqa: E402 + code_aware_split, + detect_content_type, + generate_contextual_prefix, + prepend_chunk_header, +) # --contextual flag: generate LLM context prefix per chunk (opt-in, costs ~$5-10) _CONTEXTUAL_MODE: bool = "--contextual" in sys.argv @@ -124,9 +125,7 @@ async def fetch_page_markdown(client: httpx.AsyncClient, url: str) -> tuple[str, return title, markdown -def split_into_sections( - markdown: str, url: str, page_title: str -) -> list[dict]: +def split_into_sections(markdown: str, url: str, page_title: str) -> list[dict]: """Split markdown into chunks with headers, using code-aware chunking. Returns dicts with keys: url, title, content, chunk_index, content_type. @@ -146,7 +145,7 @@ def split_into_sections( current_content: list[str] = [] for line in lines: - if line.startswith("## ") or line.startswith("### "): + if line.startswith(("## ", "### ")): if current_section and current_content: content = "\n".join(current_content).strip() if content: @@ -211,9 +210,7 @@ async def generate_embedding( """Generate embedding with concurrency limit.""" async with sem: try: - resp = await openai_client.embeddings.create( - input=text, model=EMBEDDING_MODEL - ) + resp = await openai_client.embeddings.create(input=text, model=EMBEDDING_MODEL) return resp.data[0].embedding except Exception as exc: logger.warning("Embedding error: %s", exc) @@ -257,7 +254,7 @@ async def process_and_store( contextual_prefix = prefix # Enriched content used for embedding: context marker + chunk embed_content = f"[Context: {prefix}]\n\n{raw_chunk}" - except Exception as exc: # noqa: BLE001 + except Exception as exc: logger.warning("Contextual prefix failed for %s: %s", url, exc) # Fall back to header-enriched content embed_content = section["content"] @@ -360,7 +357,12 @@ async def run_recrawl() -> None: continue chunks = await process_and_store( - pool, openai_client, sem, url, title, markdown, + pool, + openai_client, + sem, + url, + title, + markdown, contextual_client=contextual_client, ) total_chunks += chunks diff --git a/pyproject.toml b/pyproject.toml index 1451153..79a3496 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -42,5 +42,20 @@ Repository = "https://github.com/FaustoS88/Pydantic-AI-Pinescript-Expert" line-length = 100 target-version = "py311" +[tool.ruff.lint] +select = [ + "E", # pycodestyle errors + "W", # pycodestyle warnings + "F", # pyflakes + "I", # isort + "UP", # pyupgrade + "B", # flake8-bugbear + "C4", # flake8-comprehensions + "SIM", # flake8-simplify + "PIE", # flake8-pie + "DTZ", # flake8-datetimez + "RUF", # ruff-specific +] + [tool.pytest.ini_options] asyncio_mode = "auto" diff --git a/rag_utils.py b/rag_utils.py index 7f95667..80f921f 100644 --- a/rag_utils.py +++ b/rag_utils.py @@ -14,8 +14,8 @@ from __future__ import annotations import logging +from collections.abc import Sequence from dataclasses import dataclass, field -from typing import Sequence import numpy as np @@ -46,7 +46,7 @@ class RetrievedDoc: url: str title: str content: str - vector_score: float = 0.0 # cosine similarity (0–1) + vector_score: float = 0.0 # cosine similarity (0-1) bm25_score: float = 0.0 # ts_rank score rrf_score: float = 0.0 # fused score rerank_score: float = 0.0 # cross-encoder score @@ -181,7 +181,7 @@ def rerank_docs( pairs = [[query, doc.content] for doc in docs] scores = encoder.predict(pairs) - for doc, score in zip(docs, scores): + for doc, score in zip(docs, scores, strict=True): doc.rerank_score = float(score) reranked = sorted(docs, key=lambda d: d.rerank_score, reverse=True) @@ -270,9 +270,7 @@ def _split_recursive( # If this single part is too large, recurse with finer separators if len(part) > chunk_size and remaining_separators: - sub_chunks = _split_recursive( - part, remaining_separators, chunk_size, chunk_overlap - ) + sub_chunks = _split_recursive(part, remaining_separators, chunk_size, chunk_overlap) chunks.extend(sub_chunks) current = "" else: @@ -446,9 +444,7 @@ async def hybrid_retrieve( import pydantic_core # 1. Generate query embedding - embedding_resp = await openai_client.embeddings.create( - input=query, model=embedding_model - ) + embedding_resp = await openai_client.embeddings.create(input=query, model=embedding_model) query_embedding = embedding_resp.data[0].embedding embedding_json = pydantic_core.to_json(query_embedding).decode() @@ -649,7 +645,7 @@ async def generate_contextual_prefix( ) prefix = resp.choices[0].message.content.strip() return prefix if prefix else page_title - except Exception as exc: # noqa: BLE001 + except Exception as exc: logger.warning("generate_contextual_prefix: LLM call failed, using fallback: %s", exc) return page_title @@ -678,8 +674,7 @@ def detect_content_type(content: str) -> str: code_blocks = content.count("```") // 2 has_param_keywords = any( - k in content.lower() - for k in ["parameter", "argument", "syntax", "returns", "return type"] + k in content.lower() for k in ["parameter", "argument", "syntax", "returns", "return type"] ) word_count = len(content.split()) diff --git a/run.py b/run.py old mode 100644 new mode 100755 index e51f6d8..735ca9d --- a/run.py +++ b/run.py @@ -4,20 +4,24 @@ This script provides a convenient way to run different commands. """ +import asyncio import os import sys -import asyncio + from dotenv import load_dotenv # Load environment variables load_dotenv() + def print_usage(): """Print usage instructions""" print("Usage:") print(" python run.py interactive [--model PRESET] - Run the interactive shell") print(" python run.py query [--model PRESET] - Run a single query") - print(" python run.py populate - Populate the database with documentation") + print( + " python run.py populate - Populate the database with documentation" + ) print(" python run.py check - Check the database setup") print() print("Model presets (via --model):") @@ -27,35 +31,43 @@ def print_usage(): print(" flash - google/gemini-3-flash-preview") print(" - pass a raw OpenRouter model ID directly") + async def check_database(): """Check the database setup""" from agent import database_connect from db_schema import validate_schema - + try: async with database_connect(False) as pool: is_valid = await validate_schema(pool) if is_valid: print("Database schema is valid") - + # Count documents count = await pool.fetchval("SELECT COUNT(*) FROM pinescript_docs") print(f"Database contains {count} documentation sections") - + if count == 0: - print("Warning: No documentation found in database. Run 'python run.py populate ' to add documentation.") + print( + "Warning: No documentation found in database. " + "Run 'python run.py populate ' to add documentation." + ) else: - print("Database schema is not valid. Run 'python run.py populate ' to set up the database.") + print( + "Database schema is not valid. " + "Run 'python run.py populate ' to set up the database." + ) except Exception as e: print(f"Error checking database: {e}") + def _extract_model_flag(args: list[str]) -> tuple[str | None, list[str]]: """Pull --model VALUE out of args, return (preset, remaining_args).""" if "--model" in args: idx = args.index("--model") if idx + 1 < len(args): preset = args[idx + 1] - remaining = args[:idx] + args[idx + 2:] + remaining = args[:idx] + args[idx + 2 :] return preset, remaining return None, args return None, args @@ -72,6 +84,7 @@ async def run_command(command, args): if command == "interactive": from interactive import main + await main(preset=preset) elif command == "query": @@ -81,48 +94,57 @@ async def run_command(command, args): return from agent import run_agent + query = " ".join(args) if preset: from config import MODEL_PRESETS, get_preset - label = preset if preset not in MODEL_PRESETS else f"{preset} ({get_preset(preset)['model']})" + + label = ( + preset + if preset not in MODEL_PRESETS + else f"{preset} ({get_preset(preset)['model']})" + ) print(f"Model: {label}") result = await run_agent(query, preset=preset) if result: print("\nResponse:") print(result.output.response) - + elif command == "populate": if not args: print("Error: No docs directory provided") print("Usage: python run.py populate ") return - + docs_dir = args[0] if not os.path.exists(docs_dir): print(f"Error: Directory {docs_dir} does not exist") return - + from populate_db import build_search_db + await build_search_db(docs_dir) - + elif command == "check": await check_database() - + else: print(f"Unknown command: {command}") print_usage() + def main(): """Main entry point""" if len(sys.argv) < 2: print_usage() return - + command = sys.argv[1] args = sys.argv[2:] - + asyncio.run(run_command(command, args)) + if __name__ == "__main__": main() diff --git a/setup.py b/setup.py old mode 100644 new mode 100755 index 30e2b50..e4321a0 --- a/setup.py +++ b/setup.py @@ -4,67 +4,70 @@ This script helps set up the environment and database for the agent. """ -import sys -import subprocess import shutil +import subprocess +import sys from pathlib import Path + def create_virtualenv(): """Create a virtual environment and install dependencies""" print("Creating virtual environment...") venv_path = Path("venv") - + if venv_path.exists(): print("Virtual environment already exists") return True - + try: # Create virtual environment subprocess.run([sys.executable, "-m", "venv", "venv"], check=True) - + # Get the pip path if sys.platform == "win32": pip_path = venv_path / "Scripts" / "pip" else: pip_path = venv_path / "bin" / "pip" - + # Upgrade pip subprocess.run([str(pip_path), "install", "--upgrade", "pip"], check=True) - + # Install dependencies subprocess.run([str(pip_path), "install", "-r", "requirements.txt"], check=True) - + print("Virtual environment created and dependencies installed") return True except Exception as e: print(f"Error creating virtual environment: {e}") return False + def create_env_file(): """Create a .env file if it doesn't exist""" env_path = Path(".env") env_example_path = Path(".env.example") - + if env_path.exists(): print(".env file already exists") return - + if not env_example_path.exists(): print("Error: .env.example file not found") return - + shutil.copy(env_example_path, env_path) print(".env file created from .env.example") print("Please edit the .env file to add your API keys and database URL") + def make_script_executable(): """Make the shell script executable""" script_path = Path("pinescript-agent.sh") - + if not script_path.exists(): print("Error: pinescript-agent.sh not found") return - + try: script_path.chmod(0o755) print("Made pinescript-agent.sh executable") @@ -72,19 +75,20 @@ def make_script_executable(): print(f"Error making script executable: {e}") print("You may need to run: chmod +x pinescript-agent.sh") + def main(): """Main entry point""" print("Setting up Pine Script Expert Agent...") - + # Create virtual environment create_virtualenv() - + # Create .env file create_env_file() - + # Make script executable make_script_executable() - + print("\nSetup complete!") print("\nNext steps:") print("1. Edit the .env file to add your API keys and database URL") @@ -92,5 +96,6 @@ def main(): print("3. Populate the database: python run.py populate ") print("4. Start the interactive shell: python run.py interactive") + if __name__ == "__main__": main() diff --git a/streamlit_ui.py b/streamlit_ui.py index e79c9f1..2546a42 100644 --- a/streamlit_ui.py +++ b/streamlit_ui.py @@ -5,22 +5,23 @@ messages, so the model sees its own prior answers in multi-turn chats. """ -import streamlit as st import asyncio import os -import sys import pickle +import sys from pathlib import Path + +import streamlit as st from dotenv import load_dotenv from openai import AsyncOpenAI # v2 message types for proper history reconstruction -from pydantic_ai.messages import ModelRequest, ModelResponse, UserPromptPart, TextPart +from pydantic_ai.messages import ModelRequest, ModelResponse, TextPart, UserPromptPart current_dir = Path(__file__).parent sys.path.append(str(current_dir)) -from agent import database_connect, Dependencies, pinescript_agent # noqa: E402 +from agent import Dependencies, database_connect, pinescript_agent # noqa: E402 from db_schema import validate_schema # noqa: E402 load_dotenv(override=True) @@ -32,6 +33,7 @@ # Chat history persistence (pickle — swap for Harness Memory later) # --------------------------------------------------------------------------- + def load_chat_history() -> list[dict]: if os.path.exists(HISTORY_FILE): try: @@ -60,18 +62,15 @@ def save_chat_history(messages: list[dict]) -> None: # ModelResponse for assistant turns # --------------------------------------------------------------------------- + def build_message_history(messages: list[dict]) -> list: """Convert Streamlit chat dicts to Pydantic AI model messages.""" model_messages = [] for msg in messages: if msg["role"] == "user": - model_messages.append( - ModelRequest(parts=[UserPromptPart(content=msg["content"])]) - ) + model_messages.append(ModelRequest(parts=[UserPromptPart(content=msg["content"])])) elif msg["role"] == "assistant": - model_messages.append( - ModelResponse(parts=[TextPart(content=msg["content"])]) - ) + model_messages.append(ModelResponse(parts=[TextPart(content=msg["content"])])) return model_messages @@ -98,6 +97,7 @@ def build_message_history(messages: list[dict]) -> list: # Setup verification # --------------------------------------------------------------------------- + async def verify_setup(): openai_api_key = os.getenv("OPENAI_API_KEY") if not openai_api_key or openai_api_key in ("YOUR_OPENAI_API_KEY", "sk-...", ""): @@ -120,6 +120,7 @@ async def verify_setup(): # Query processing — v2: passes full history (user + assistant) to the model # --------------------------------------------------------------------------- + async def process_query(prompt: str, history: list[dict]): openai_client = AsyncOpenAI(api_key=os.getenv("OPENAI_API_KEY")) @@ -222,4 +223,4 @@ async def process_query(prompt: str, history: list[dict]): st.session_state.messages.append( {"role": "assistant", "content": f"Error: {e}", "snippets": 0} ) - save_chat_history(st.session_state.messages) \ No newline at end of file + save_chat_history(st.session_state.messages) diff --git a/tests/ragas_eval.py b/tests/ragas_eval.py index 032b73c..d1da595 100644 --- a/tests/ragas_eval.py +++ b/tests/ragas_eval.py @@ -1,7 +1,8 @@ """RAGAS evaluation harness for the PineScript Expert RAG pipeline. Usage: - python tests/ragas_eval.py [--output results/baseline_YYYYMMDD.json] [--retrieval baseline|tier1] + python tests/ragas_eval.py [--output results/baseline_YYYYMMDD.json] + [--retrieval baseline|tier1] --retrieval baseline L2 vector search, no threshold, no reranking (default) --retrieval tier1 Hybrid cosine+BM25, RRF, similarity threshold, cross-encoder rerank, MMR @@ -22,7 +23,7 @@ import logging import os import sys -from datetime import datetime, timezone +from datetime import UTC, datetime from pathlib import Path from typing import Any @@ -51,12 +52,12 @@ from config import ( # noqa: E402 DEFAULT_DATABASE_URL, EMBEDDING_MODEL, - VECTOR_SEARCH_LIMIT, HYBRID_SEARCH_ALPHA, - SIMILARITY_THRESHOLD, - RETRIEVAL_CANDIDATES, - RERANK_TOP_N, MMR_LAMBDA, + RERANK_TOP_N, + RETRIEVAL_CANDIDATES, + SIMILARITY_THRESHOLD, + VECTOR_SEARCH_LIMIT, ) from rag_utils import hybrid_retrieve # noqa: E402 @@ -79,6 +80,7 @@ # Retrieval — baseline (L2) or tier1 (hybrid cosine+BM25, rerank, MMR) # --------------------------------------------------------------------------- + async def retrieve_contexts_baseline( question: str, openai_client: AsyncOpenAI, @@ -94,16 +96,12 @@ async def retrieve_contexts_baseline( embedding_json = pydantic_core.to_json(embedding_vector).decode() rows = await pool.fetch( - f"SELECT url, title, content FROM pinescript_docs " - f"ORDER BY embedding <-> $1 LIMIT {limit}", + f"SELECT url, title, content FROM pinescript_docs ORDER BY embedding <-> $1 LIMIT {limit}", embedding_json, ) if not rows: return [] - return [ - f"# {row['title']}\nURL: {row['url']}\n\n{row['content']}" - for row in rows - ] + return [f"# {row['title']}\nURL: {row['url']}\n\n{row['content']}" for row in rows] async def retrieve_contexts_tier1( @@ -131,16 +129,14 @@ async def retrieve_contexts_tier1( ) if not docs: return [] - return [ - f"# {doc.title}\nURL: {doc.url}\n\n{doc.content}" - for doc in docs - ] + return [f"# {doc.title}\nURL: {doc.url}\n\n{doc.content}" for doc in docs] # --------------------------------------------------------------------------- # Generation — direct OpenRouter call (bypasses pydantic-ai version issues) # --------------------------------------------------------------------------- + async def generate_answer( question: str, contexts: list[str], @@ -153,10 +149,7 @@ async def generate_answer( context_block = "\n\n---\n\n".join(contexts) if contexts else "(no context retrieved)" - user_msg = ( - f"## Retrieved Documentation\n\n{context_block}\n\n" - f"---\n\n## Question\n\n{question}" - ) + user_msg = f"## Retrieved Documentation\n\n{context_block}\n\n---\n\n## Question\n\n{question}" try: resp = await http_client.post( @@ -179,7 +172,7 @@ async def generate_answer( resp.raise_for_status() data = resp.json() return data["choices"][0]["message"]["content"] - except Exception as exc: # noqa: BLE001 + except Exception as exc: logger.error("Generation error for question %r: %s", question[:60], exc) return f"[ERROR] {exc}" @@ -188,6 +181,7 @@ async def generate_answer( # RAGAS evaluation (v0.4 API) # --------------------------------------------------------------------------- + def run_ragas( questions: list[str], answers: list[str], @@ -202,13 +196,13 @@ def run_ragas( Both are old-style metrics compatible with evaluate(). """ try: - from ragas import evaluate, RunConfig - from ragas.dataset_schema import SingleTurnSample, EvaluationDataset - from ragas.metrics._faithfulness import Faithfulness - from ragas.metrics._nv_metrics import ContextRelevance - from ragas.llms import llm_factory, LangchainLLMWrapper from langchain_openai import ChatOpenAI from openai import OpenAI + from ragas import RunConfig, evaluate + from ragas.dataset_schema import EvaluationDataset, SingleTurnSample + from ragas.llms import LangchainLLMWrapper, llm_factory + from ragas.metrics._faithfulness import Faithfulness + from ragas.metrics._nv_metrics import ContextRelevance except ImportError as exc: logger.error( "RAGAS dependencies missing. Run: pip install ragas datasets langchain-openai\n%s", @@ -224,9 +218,7 @@ def run_ragas( instructor_llm.model_args["max_tokens"] = 8192 # LangchainLLMWrapper for ContextRelevance (uses agenerate_text interface) - langchain_llm = LangchainLLMWrapper( - ChatOpenAI(model="gpt-4.1-mini", api_key=openai_api_key) - ) + langchain_llm = LangchainLLMWrapper(ChatOpenAI(model="gpt-4.1-mini", api_key=openai_api_key)) metrics = [ Faithfulness(llm=instructor_llm), @@ -235,24 +227,28 @@ def run_ragas( # Build evaluation dataset samples = [] - for q, a, ctx in zip(questions, answers, contexts): - samples.append(SingleTurnSample( - user_input=q, - response=a, - retrieved_contexts=ctx, - )) + for q, a, ctx in zip(questions, answers, contexts, strict=True): + samples.append( + SingleTurnSample( + user_input=q, + response=a, + retrieved_contexts=ctx, + ) + ) dataset = EvaluationDataset(samples=samples) run_config = RunConfig(max_retries=2, max_wait=120, max_workers=4) - logger.info("Running RAGAS evaluate() on %d samples (2 metrics x %d) …", len(questions), len(questions)) + logger.info( + "Running RAGAS evaluate() on %d samples (2 metrics x %d) …", len(questions), len(questions) + ) result = evaluate(dataset=dataset, metrics=metrics, run_config=run_config) # Convert to plain dict try: scores_df = result.to_pandas() return scores_df.to_dict(orient="list") - except Exception: # noqa: BLE001 + except Exception: return dict(result) @@ -348,7 +344,7 @@ def print_report(report: dict[str, Any], meta: dict[str, Any]) -> None: f_ = it.get("faithfulness") cp_s = f"{cp:.2f}" if cp is not None else "n/a" f_s = f"{f_:.2f}" if f_ is not None else "n/a" - print(f" Q: \"{it['question'][:70]}\" → CP={cp_s}, F={f_s}") + print(f' Q: "{it["question"][:70]}" → CP={cp_s}, F={f_s}') print(sep) @@ -397,7 +393,9 @@ async def run_evaluation(output_path: Path, retrieval_mode: str = "baseline") -> contexts: list[list[str]] = [] items: list[dict[str, Any]] = [] - retrieve_fn = retrieve_contexts_tier1 if retrieval_mode == "tier1" else retrieve_contexts_baseline + retrieve_fn = ( + retrieve_contexts_tier1 if retrieval_mode == "tier1" else retrieve_contexts_baseline + ) # Partial save: written after every question so we never lose progress on crash RESULTS_DIR.mkdir(exist_ok=True) @@ -411,7 +409,7 @@ def _save_partial() -> None: "contexts": c, "item": item, } - for q, a, c, item in zip(questions, answers, contexts, items) + for q, a, c, item in zip(questions, answers, contexts, items, strict=True) ] with open(partial_path, "w") as fh: json.dump(partial_data, fh, indent=2, default=str) @@ -449,7 +447,7 @@ def _save_partial() -> None: # RAGAS evaluation ragas_scores = run_ragas(questions, answers, contexts) - run_date = datetime.now(timezone.utc).strftime("%Y-%m-%d") + run_date = datetime.now(UTC).strftime("%Y-%m-%d") meta = { "date": run_date, "questions": len(testset), @@ -480,7 +478,9 @@ def _save_partial() -> None: def main() -> None: - parser = argparse.ArgumentParser(description="Run RAGAS evaluation on the PineScript Expert RAG pipeline.") + parser = argparse.ArgumentParser( + description="Run RAGAS evaluation on the PineScript Expert RAG pipeline." + ) parser.add_argument( "--retrieval", choices=["baseline", "tier1"], @@ -497,7 +497,7 @@ def main() -> None: if args.output is None: tag = args.retrieval - args.output = RESULTS_DIR / f"{tag}_{datetime.now(timezone.utc).strftime('%Y%m%d')}.json" + args.output = RESULTS_DIR / f"{tag}_{datetime.now(UTC).strftime('%Y%m%d')}.json" asyncio.run(run_evaluation(args.output, retrieval_mode=args.retrieval)) diff --git a/tests/test_agent.py b/tests/test_agent.py index c1a9486..5c41891 100644 --- a/tests/test_agent.py +++ b/tests/test_agent.py @@ -1,166 +1,171 @@ -"""Tests for the RAG retrieve tool — mocked DB + embeddings, no container needed.""" +"""Tests for the RAG retrieve tool — mocked retrieval pipeline, no container needed.""" from __future__ import annotations -from unittest.mock import AsyncMock, MagicMock +from unittest.mock import AsyncMock, MagicMock, patch import pytest -from agent import Dependencies, retrieve, pinescript_agent - +from agent import ( + Dependencies, + PineScriptResult, + pinescript_agent, + pinescript_rag, + resolve_model, + retrieve, +) +from config import EMBEDDING_MODEL, OPENROUTER_DEFAULT_MODEL +from rag_utils import RetrievedDoc # --------------------------------------------------------------------------- # Fixtures # --------------------------------------------------------------------------- -def _make_deps(rows=None, use_openrouter=False): - """Build a Dependencies object with mocked OpenAI client and asyncpg pool.""" - # Mock OpenAI embeddings response - embedding_obj = MagicMock() - embedding_obj.data = [MagicMock(embedding=[0.1] * 1536)] - - openai_mock = AsyncMock() - openai_mock.embeddings.create = AsyncMock(return_value=embedding_obj) - - # Mock asyncpg pool - pool_mock = AsyncMock() - pool_mock.fetch = AsyncMock(return_value=rows or []) - return Dependencies( - openai=openai_mock, - pool=pool_mock, - openrouter_api_key="sk-or-test" if use_openrouter else None, - use_openrouter=use_openrouter, - ) +def _make_deps(): + """Build a Dependencies object with mocked OpenAI client and asyncpg pool.""" + return Dependencies(openai=AsyncMock(), pool=AsyncMock()) def _make_ctx(deps): """Build a minimal RunContext-like object for the retrieve tool.""" ctx = MagicMock() ctx.deps = deps - ctx.custom_data = {} return ctx +def _doc(url: str, title: str, content: str) -> RetrievedDoc: + return RetrievedDoc(url=url, title=title, content=content) + + # --------------------------------------------------------------------------- # retrieve() tool tests +# +# retrieve() delegates the whole vector + BM25 → RRF → threshold → rerank → +# MMR pipeline to rag_utils.hybrid_retrieve (covered in test_rag_improvements), +# so these tests patch it and assert the tool's own contract: formatting, +# the empty case, and error handling. # --------------------------------------------------------------------------- + class TestRetrieveTool: @pytest.mark.asyncio async def test_returns_formatted_docs(self) -> None: - rows = [ - {"url": "https://docs.tv/plot", "title": "plot()", "content": "Plots a line on the chart."}, - {"url": "https://docs.tv/hline", "title": "hline()", "content": "Draws a horizontal line."}, + docs = [ + _doc("https://docs.tv/plot", "plot()", "Plots a line on the chart."), + _doc("https://docs.tv/hline", "hline()", "Draws a horizontal line."), ] - deps = _make_deps(rows=rows) - ctx = _make_ctx(deps) + ctx = _make_ctx(_make_deps()) - result = await retrieve(ctx, "how to plot a line") + with patch("agent.hybrid_retrieve", AsyncMock(return_value=docs)): + result = await retrieve(ctx, "how to plot a line") assert "plot()" in result assert "hline()" in result assert "Plots a line" in result - assert ctx.custom_data["snippets_used"] == 2 + assert "https://docs.tv/hline" in result @pytest.mark.asyncio async def test_returns_message_when_no_docs(self) -> None: - deps = _make_deps(rows=[]) - ctx = _make_ctx(deps) + ctx = _make_ctx(_make_deps()) - result = await retrieve(ctx, "nonexistent topic") + with patch("agent.hybrid_retrieve", AsyncMock(return_value=[])): + result = await retrieve(ctx, "nonexistent topic") assert "No relevant documentation found" in result @pytest.mark.asyncio - async def test_calls_openai_embeddings(self) -> None: - deps = _make_deps(rows=[]) - ctx = _make_ctx(deps) - - await retrieve(ctx, "test query") - - deps.openai.embeddings.create.assert_called_once() - call_kwargs = deps.openai.embeddings.create.call_args - assert call_kwargs.kwargs["input"] == "test query" - assert "embedding" in call_kwargs.kwargs["model"] - - @pytest.mark.asyncio - async def test_queries_db_with_embedding(self) -> None: - deps = _make_deps(rows=[]) - ctx = _make_ctx(deps) - - await retrieve(ctx, "indicator question") - - deps.pool.fetch.assert_called_once() - sql_arg = deps.pool.fetch.call_args[0][0] - assert "pinescript_docs" in sql_arg - assert "LIMIT" in sql_arg - - @pytest.mark.asyncio - async def test_handles_embedding_error(self) -> None: + async def test_passes_deps_and_query_to_pipeline(self) -> None: deps = _make_deps() - deps.openai.embeddings.create = AsyncMock(side_effect=Exception("API down")) ctx = _make_ctx(deps) + fake_retrieve = AsyncMock(return_value=[]) - result = await retrieve(ctx, "test") + with patch("agent.hybrid_retrieve", fake_retrieve): + await retrieve(ctx, "indicator question") - assert "Error" in result + fake_retrieve.assert_called_once() + kwargs = fake_retrieve.call_args.kwargs + assert kwargs["query"] == "indicator question" + assert kwargs["pool"] is deps.pool + assert kwargs["openai_client"] is deps.openai + assert kwargs["embedding_model"] == EMBEDDING_MODEL @pytest.mark.asyncio - async def test_handles_db_error(self) -> None: - deps = _make_deps() - ctx = _make_ctx(deps) - deps.pool.fetch = AsyncMock(side_effect=Exception("connection refused")) + async def test_handles_retrieval_error(self) -> None: + ctx = _make_ctx(_make_deps()) - result = await retrieve(ctx, "test") + with patch("agent.hybrid_retrieve", AsyncMock(side_effect=Exception("connection refused"))): + result = await retrieve(ctx, "test") assert "Error" in result + assert "connection refused" in result @pytest.mark.asyncio async def test_single_doc_formatting(self) -> None: - rows = [ - {"url": "https://docs.tv/var", "title": "Variables", "content": "Use var to declare."}, - ] - deps = _make_deps(rows=rows) - ctx = _make_ctx(deps) + docs = [_doc("https://docs.tv/var", "Variables", "Use var to declare.")] + ctx = _make_ctx(_make_deps()) - result = await retrieve(ctx, "variables") + with patch("agent.hybrid_retrieve", AsyncMock(return_value=docs)): + result = await retrieve(ctx, "variables") assert "# Variables" in result assert "https://docs.tv/var" in result - assert ctx.custom_data["snippets_used"] == 1 + assert "Use var to declare." in result # --------------------------------------------------------------------------- -# Agent wiring +# Agent wiring — v2: the tool is registered on the Capability, and the output +# schema is `output_type` (was `result_type` in v1). # --------------------------------------------------------------------------- + class TestAgentWiring: - def test_agent_has_retrieve_tool(self) -> None: - tool_names = [t.name for t in pinescript_agent._function_tools.values()] - assert "retrieve" in tool_names + def test_capability_registers_retrieve_tool(self) -> None: + assert "retrieve" in pinescript_rag.get_toolset().tools - def test_agent_result_type(self) -> None: - from agent import PineScriptResult - assert pinescript_agent.result_type == PineScriptResult + def test_agent_output_type(self) -> None: + assert pinescript_agent.output_type == PineScriptResult def test_agent_uses_config_model(self) -> None: from config import DEFAULT_MODEL + # Agent was initialised with DEFAULT_MODEL from config assert DEFAULT_MODEL == "openai:gpt-4o-mini" # --------------------------------------------------------------------------- -# Dependencies — OpenRouter routing +# Model routing — v2 moved this off Dependencies and onto resolve_model(), +# which reads OPENROUTER_API_KEY at call time. # --------------------------------------------------------------------------- -class TestOpenRouterRouting: - def test_deps_without_openrouter(self) -> None: - deps = _make_deps(use_openrouter=False) - assert deps.use_openrouter is False - assert deps.openrouter_api_key is None - def test_deps_with_openrouter(self) -> None: - deps = _make_deps(use_openrouter=True) - assert deps.use_openrouter is True - assert deps.openrouter_api_key == "sk-or-test" +class TestResolveModel: + def test_returns_none_without_openrouter_key(self, monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.delenv("OPENROUTER_API_KEY", raising=False) + assert resolve_model() is None + + def test_uses_openrouter_default_with_key(self, monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("OPENROUTER_API_KEY", "sk-or-test") + model = resolve_model() + assert model is not None + assert model.model_name == OPENROUTER_DEFAULT_MODEL + + def test_known_preset_resolves_to_preset_model(self, monkeypatch: pytest.MonkeyPatch) -> None: + from config import get_preset + + monkeypatch.setenv("OPENROUTER_API_KEY", "sk-or-test") + model = resolve_model("flash") + assert model is not None + assert model.model_name == str(get_preset("flash")["model"]) + + def test_raw_model_id_passes_through(self, monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("OPENROUTER_API_KEY", "sk-or-test") + model = resolve_model("vendor/some-raw-model") + assert model is not None + assert model.model_name == "vendor/some-raw-model" + + def test_preset_without_key_falls_back_to_default( + self, monkeypatch: pytest.MonkeyPatch + ) -> None: + monkeypatch.delenv("OPENROUTER_API_KEY", raising=False) + assert resolve_model("flash") is None diff --git a/tests/test_config.py b/tests/test_config.py index 2668073..06a4222 100644 --- a/tests/test_config.py +++ b/tests/test_config.py @@ -6,47 +6,57 @@ class TestConfigDefaults: def test_default_model(self) -> None: from config import DEFAULT_MODEL + assert DEFAULT_MODEL == "openai:gpt-4o-mini" def test_temperature_is_float(self) -> None: from config import LLM_TEMPERATURE + assert isinstance(LLM_TEMPERATURE, float) assert 0.0 <= LLM_TEMPERATURE <= 2.0 def test_max_tokens_is_int(self) -> None: from config import LLM_MAX_TOKENS + assert isinstance(LLM_MAX_TOKENS, int) assert LLM_MAX_TOKENS > 0 def test_embedding_model(self) -> None: from config import EMBEDDING_MODEL + assert EMBEDDING_MODEL == "text-embedding-3-small" def test_embedding_dimension(self) -> None: from config import EMBEDDING_DIMENSION + assert EMBEDDING_DIMENSION == 1536 def test_vector_search_limit(self) -> None: from config import VECTOR_SEARCH_LIMIT + assert isinstance(VECTOR_SEARCH_LIMIT, int) assert VECTOR_SEARCH_LIMIT > 0 def test_openrouter_base_url(self) -> None: from config import OPENROUTER_BASE_URL + assert "openrouter.ai" in OPENROUTER_BASE_URL def test_database_url_has_scheme(self) -> None: from config import DEFAULT_DATABASE_URL + assert DEFAULT_DATABASE_URL.startswith("postgresql://") class TestModelPresets: def test_default_preset_exists(self) -> None: from config import MODEL_PRESETS + assert "default" in MODEL_PRESETS def test_codex_preset(self) -> None: from config import get_preset + preset = get_preset("codex") assert preset["model"] == "openai/gpt-5.3-codex" assert preset["temperature"] == 0.1 @@ -54,22 +64,26 @@ def test_codex_preset(self) -> None: def test_opus_preset(self) -> None: from config import get_preset + preset = get_preset("opus") assert preset["model"] == "anthropic/claude-opus-4-6" assert preset["max_tokens"] == 4096 def test_flash_preset(self) -> None: from config import get_preset + preset = get_preset("flash") assert "gemini" in preset["model"] def test_unknown_preset_returns_default(self) -> None: - from config import get_preset, MODEL_PRESETS + from config import MODEL_PRESETS, get_preset + result = get_preset("nonexistent") assert result == MODEL_PRESETS["default"] def test_all_presets_have_required_keys(self) -> None: from config import MODEL_PRESETS + for name, preset in MODEL_PRESETS.items(): assert "model" in preset, f"preset '{name}' missing 'model'" assert "temperature" in preset, f"preset '{name}' missing 'temperature'" diff --git a/tests/test_models.py b/tests/test_models.py index 071b2f6..5a3ba37 100644 --- a/tests/test_models.py +++ b/tests/test_models.py @@ -2,16 +2,18 @@ from __future__ import annotations +from dataclasses import fields + import pytest from pydantic import ValidationError -from agent import PineScriptResult, Dependencies - +from agent import Dependencies, PineScriptResult # --------------------------------------------------------------------------- # PineScriptResult — validation # --------------------------------------------------------------------------- + class TestPineScriptResult: def test_valid_result(self) -> None: result = PineScriptResult( @@ -24,9 +26,7 @@ def test_valid_result(self) -> None: assert result.snippets_used == 3 def test_zero_snippets(self) -> None: - result = PineScriptResult( - query="test", response="answer", snippets_used=0 - ) + result = PineScriptResult(query="test", response="answer", snippets_used=0) assert result.snippets_used == 0 def test_missing_query_raises(self) -> None: @@ -51,23 +51,17 @@ def test_empty_strings_allowed(self) -> None: def test_long_response(self) -> None: long_text = "x" * 50_000 - result = PineScriptResult( - query="test", response=long_text, snippets_used=1 - ) + result = PineScriptResult(query="test", response=long_text, snippets_used=1) assert len(result.response) == 50_000 def test_serialisation_roundtrip(self) -> None: - original = PineScriptResult( - query="plot question", response="use plot()", snippets_used=5 - ) + original = PineScriptResult(query="plot question", response="use plot()", snippets_used=5) data = original.model_dump() restored = PineScriptResult(**data) assert restored == original def test_json_roundtrip(self) -> None: - original = PineScriptResult( - query="test", response="answer", snippets_used=2 - ) + original = PineScriptResult(query="test", response="answer", snippets_used=2) json_str = original.model_dump_json() restored = PineScriptResult.model_validate_json(json_str) assert restored == original @@ -77,21 +71,22 @@ def test_json_roundtrip(self) -> None: # Dependencies — dataclass init # --------------------------------------------------------------------------- + class TestDependencies: - def test_default_values(self) -> None: - deps = Dependencies(openai=None, pool=None) # type: ignore[arg-type] - assert deps.openrouter_api_key is None - assert deps.use_openrouter is False - - def test_with_openrouter(self) -> None: - deps = Dependencies( - openai=None, # type: ignore[arg-type] - pool=None, # type: ignore[arg-type] - openrouter_api_key="sk-or-test", - use_openrouter=True, - ) - assert deps.openrouter_api_key == "sk-or-test" - assert deps.use_openrouter is True + def test_only_openai_and_pool(self) -> None: + """v2 trimmed Dependencies to the two fields tools actually read.""" + field_names = {f.name for f in fields(Dependencies)} + assert field_names == {"openai", "pool"} + + def test_rejects_removed_openrouter_fields(self) -> None: + """Model routing moved to resolve_model(); these fields no longer exist.""" + with pytest.raises(TypeError): + Dependencies( + openai=None, # type: ignore[arg-type] + pool=None, # type: ignore[arg-type] + openrouter_api_key="sk-or-test", # type: ignore[call-arg] + use_openrouter=True, # type: ignore[call-arg] + ) def test_fields_accessible(self) -> None: deps = Dependencies(openai="mock_client", pool="mock_pool") # type: ignore[arg-type] diff --git a/tests/test_rag_improvements.py b/tests/test_rag_improvements.py index 08aaf19..f70c060 100644 --- a/tests/test_rag_improvements.py +++ b/tests/test_rag_improvements.py @@ -19,15 +19,14 @@ RetrievedDoc, _cosine_sim, apply_similarity_threshold, + hybrid_retrieve, mmr_select, prepend_chunk_header, reciprocal_rank_fusion, recursive_character_split, rerank_docs, - hybrid_retrieve, ) - # --------------------------------------------------------------------------- # Helpers # --------------------------------------------------------------------------- @@ -59,7 +58,7 @@ def _doc( class TestContextualChunkHeaders: def test_prepends_header_with_different_section(self): result = prepend_chunk_header("Arrays", "Sorting", "Use array.sort().") - assert result.startswith("Document: PineScript v5 Reference | Section: Arrays > Sorting") + assert result.startswith("Document: PineScript v6 Reference | Section: Arrays > Sorting") assert "Use array.sort()." in result def test_prepends_header_same_section_as_page(self): @@ -74,7 +73,7 @@ def test_content_preserved(self): def test_empty_content(self): result = prepend_chunk_header("Page", "Section", "") - assert "Document: PineScript v5 Reference" in result + assert "Document: PineScript v6 Reference" in result # --------------------------------------------------------------------------- @@ -412,13 +411,16 @@ def _get_crawler(): def test_split_produces_chunks_with_headers(self): crawler = self._get_crawler() - markdown = "# Arrays\n\n## Sorting\n\nUse array.sort() to sort arrays.\n\n## Inserting\n\nUse array.push()." + markdown = ( + "# Arrays\n\n## Sorting\n\nUse array.sort() to sort arrays.\n\n" + "## Inserting\n\nUse array.push()." + ) sections = crawler.split_into_sections(markdown, "https://docs.tv/arrays") assert len(sections) >= 2 # Each section should have contextual header for sec in sections: - assert "Document: PineScript v5 Reference" in sec["content"] + assert "Document: PineScript v6 Reference" in sec["content"] assert "Section:" in sec["content"] def test_split_long_section_into_multiple_chunks(self): diff --git a/tests/test_rag_tier2.py b/tests/test_rag_tier2.py index 5bfda7a..30eaaae 100644 --- a/tests/test_rag_tier2.py +++ b/tests/test_rag_tier2.py @@ -9,17 +9,17 @@ from __future__ import annotations -import sys import os -import pytest +import sys from unittest.mock import AsyncMock, MagicMock +import pytest + # Make project root importable sys.path.insert(0, os.path.dirname(os.path.dirname(__file__))) +from db_schema import DB_MIGRATION, DB_SCHEMA from rag_utils import code_aware_split, detect_content_type, generate_contextual_prefix -from db_schema import DB_SCHEMA, DB_MIGRATION - # =========================================================================== # TestCodeAwareSplit @@ -86,15 +86,11 @@ def test_mixed_prose_and_code(self): for chunk in result: backtick_count = chunk.count("```") # Even count means balanced fences (or zero) - assert backtick_count % 2 == 0, ( - f"Unbalanced backticks in chunk: {chunk!r}" - ) + assert backtick_count % 2 == 0, f"Unbalanced backticks in chunk: {chunk!r}" def test_inline_code_not_protected(self): """Single-backtick inline code is NOT treated as a fenced block.""" - long_text = ( - "The `ta.ema()` function computes the exponential moving average. " * 30 - ) + long_text = "The `ta.ema()` function computes the exponential moving average. " * 30 result = code_aware_split(long_text, chunk_size=200, chunk_overlap=0) # Should split into multiple chunks (inline code does not block splitting) assert len(result) > 1 @@ -106,9 +102,7 @@ def test_multiple_code_blocks(self): result = code_aware_split(text, chunk_size=80, chunk_overlap=0) for chunk in result: backtick_count = chunk.count("```") - assert backtick_count % 2 == 0, ( - f"Unbalanced backticks in chunk: {chunk!r}" - ) + assert backtick_count % 2 == 0, f"Unbalanced backticks in chunk: {chunk!r}" # Each code block must appear in its entirety somewhere for i in range(5): assert any(f"block_{i} = {i}" in c for c in result) @@ -175,7 +169,9 @@ async def test_truncates_large_page_content(self): # We check the LLM was called exactly once (no crash). client.chat.completions.create.assert_awaited_once() call_kwargs = client.chat.completions.create.call_args - prompt_text = call_kwargs.kwargs.get("messages", call_kwargs.args[0] if call_kwargs.args else []) + prompt_text = call_kwargs.kwargs.get( + "messages", call_kwargs.args[0] if call_kwargs.args else [] + ) # Extract prompt content from messages list if isinstance(prompt_text, list): full_prompt = " ".join(m.get("content", "") for m in prompt_text)