237 lines
8.8 KiB
Python
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
|