"""Base extractor for Claude-powered intelligence extraction.""" import json import logging from abc import ABC, abstractmethod import datetime from typing import Optional import anthropic from storage.db import MosaicDB from storage.models import TokenUsage from utils.api_limiter import APIRateLimiter, CircuitBreakerOpen from utils.text_utils import chunk_text, estimate_tokens logger = logging.getLogger(__name__) class BaseExtractor(ABC): """Abstract base for all Claude-powered extractors.""" EXTRACTOR_NAME = "base" # Override in subclasses def __init__(self, db: MosaicDB, api_key: str, model: str = "claude-sonnet-4-20250514", max_tokens: int = 4096, rate_limiter: Optional[APIRateLimiter] = None): self.db = db self.client = anthropic.Anthropic(api_key=api_key) self.model = model self.max_tokens = max_tokens self.rate_limiter = rate_limiter or APIRateLimiter() @abstractmethod def get_system_prompt(self) -> str: """Return the system prompt for this extractor.""" pass @abstractmethod def get_extraction_prompt(self, text: str) -> str: """Return the user prompt with document content wrapped in XML tags.""" pass @abstractmethod def get_json_schema(self) -> dict: """Return the expected JSON schema for validation.""" pass @abstractmethod def process_response(self, data: dict, doc_id: str) -> int: """Process validated extraction response and store results. Returns count of items extracted.""" pass def extract_from_document(self, doc_id: str, text: str, doc_type: str = "text", source: str = "", metadata: Optional[dict] = None) -> int: """Extract intelligence from a single document. Handles chunking, API calls, validation, and storage. Returns total count of items extracted. """ # Set status to IN_PROGRESS self.db.set_extraction_status(doc_id, self.EXTRACTOR_NAME, "IN_PROGRESS") try: # Chunk the document chunks = chunk_text(text, doc_type=doc_type, metadata=metadata) total_extracted = 0 total_input_tokens = 0 total_output_tokens = 0 for chunk in chunks: try: # Rate limit self.rate_limiter.wait_if_needed() # Build prompt with XML-wrapped content user_prompt = self.get_extraction_prompt(chunk['text']) # Call Claude response = self.client.messages.create( model=self.model, max_tokens=self.max_tokens, system=self.get_system_prompt(), messages=[{"role": "user", "content": user_prompt}] ) # Track tokens input_tokens = response.usage.input_tokens output_tokens = response.usage.output_tokens total_input_tokens += input_tokens total_output_tokens += output_tokens self.rate_limiter.record_success() # Extract text from response response_text = response.content[0].text.strip() # Parse JSON from response data = self._parse_json_response(response_text) if data is None: logger.warning("Failed to parse JSON from response for doc %s chunk %d", doc_id, chunk.get('chunk_index', 0)) continue # Validate against schema if not self._validate_response(data): logger.warning("Schema validation failed for doc %s chunk %d", doc_id, chunk.get('chunk_index', 0)) continue # Process and store count = self.process_response(data, doc_id) total_extracted += count except CircuitBreakerOpen: logger.error("Circuit breaker open — halting extraction") self.db.set_extraction_status( doc_id, self.EXTRACTOR_NAME, "FAILED", error="Circuit breaker open" ) return total_extracted except anthropic.APIError as e: self.rate_limiter.record_failure(str(e)) logger.error("API error on doc %s: %s", doc_id, e) continue except Exception as e: logger.error("Extraction error on doc %s chunk %d: %s", doc_id, chunk.get('chunk_index', 0), e) continue # Log token usage self.db.log_token_usage(TokenUsage( timestamp=datetime.datetime.now(datetime.timezone.utc).isoformat(), source=source, doc_id=doc_id, extractor=self.EXTRACTOR_NAME, input_tokens=total_input_tokens, output_tokens=total_output_tokens, model=self.model, )) # Update status self.db.set_extraction_status( doc_id, self.EXTRACTOR_NAME, "DONE", token_count=total_input_tokens + total_output_tokens ) self.db.update_document_extraction_time(doc_id) return total_extracted except Exception as e: logger.error("Extraction failed for doc %s: %s", doc_id, e) self.db.set_extraction_status( doc_id, self.EXTRACTOR_NAME, "FAILED", error=str(e) ) return 0 def estimate_tokens(self, text: str, doc_type: str = "text", metadata: Optional[dict] = None) -> dict: """Estimate token cost without calling the API (dry-run).""" chunks = chunk_text(text, doc_type=doc_type, metadata=metadata) total_input = 0 system_tokens = estimate_tokens(self.get_system_prompt()) for chunk in chunks: prompt = self.get_extraction_prompt(chunk['text']) total_input += estimate_tokens(prompt) + system_tokens return { 'chunks': len(chunks), 'estimated_input_tokens': total_input, 'estimated_output_tokens': len(chunks) * 2000, # rough estimate } def _parse_json_response(self, text: str) -> Optional[dict]: """Parse JSON from Claude response, handling markdown code blocks.""" # Try direct parse try: return json.loads(text) except json.JSONDecodeError: pass # Try extracting from markdown code block import re match = re.search(r'```(?:json)?\s*\n(.*?)\n```', text, re.DOTALL) if match: try: return json.loads(match.group(1)) except json.JSONDecodeError: pass # Try finding JSON object/array in text for start_char, end_char in [('{', '}'), ('[', ']')]: start = text.find(start_char) if start >= 0: # Find matching end depth = 0 for i in range(start, len(text)): if text[i] == start_char: depth += 1 elif text[i] == end_char: depth -= 1 if depth == 0: try: return json.loads(text[start:i+1]) except json.JSONDecodeError: break logger.warning("Could not parse JSON from response: %s...", text[:200]) return None def _validate_response(self, data: dict) -> bool: """Validate response against expected schema. Reject unexpected fields.""" schema = self.get_json_schema() if not schema: return True required_keys = schema.get('required', []) allowed_keys = set(schema.get('properties', {}).keys()) if isinstance(data, dict): # Check for unexpected top-level keys if 'items' in schema.get('properties', {}): # Wrapper object with 'items' array if 'items' not in data: logger.warning("Response missing 'items' key") return False unexpected = set(data.keys()) - allowed_keys if unexpected: logger.warning("Unexpected fields in response: %s", unexpected) return False return True elif isinstance(data, list): return True return False