Files
mosaic/extractors/base.py
T

237 lines
8.8 KiB
Python

"""Base extractor for Claude-powered intelligence extraction."""
import json
import logging
from abc import ABC, abstractmethod
from datetime 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.utcnow().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