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
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