Coverage for src / quber / db / importer.py: 89%
114 statements
« prev ^ index » next coverage.py v7.14.0, created at 2026-09-23 22:14 -0400
« prev ^ index » next coverage.py v7.14.0, created at 2026-09-23 22:14 -0400
1"""Import utility for loading JSON analyses into the database."""
3import json
4from datetime import datetime
5from pathlib import Path
6from typing import Any
8from loguru import logger
9from sqlalchemy.orm import Session
11from quber.db.connection import get_session
12from quber.db.embeddings import EmbeddingService, get_embedding_service
13from quber.db.models import Document, ExtractedTable
16def import_json_file(
17 json_path: Path | str,
18 session: Session | None = None,
19 embedding_service: EmbeddingService | None = None,
20 generate_embeddings: bool = False,
21) -> tuple[Document, int]:
22 """
23 Import a single JSON analysis file into the database.
25 Args:
26 json_path: Path to the JSON file.
27 session: Optional database session. If None, creates a new session.
28 embedding_service: Optional embedding service. If None and generate_embeddings=True, creates one.
29 generate_embeddings: Whether to generate embeddings during import.
31 Returns:
32 Tuple of (Document instance, number of tables imported).
34 Raises:
35 FileNotFoundError: If the JSON file doesn't exist.
36 ValueError: If the JSON structure is invalid.
37 """
38 json_path = Path(json_path)
39 if not json_path.exists():
40 raise FileNotFoundError(f"JSON file not found: {json_path}")
42 logger.info(f"Importing {json_path.name}...")
44 with open(json_path) as f:
45 data: dict[str, Any] = json.load(f)
47 # Parse extraction date
48 extraction_date_str = data.get("extraction_date", "")
49 try:
50 extraction_date = datetime.fromisoformat(extraction_date_str)
51 except ValueError:
52 logger.warning(f"Invalid extraction_date format: {extraction_date_str}, using current time")
53 extraction_date = datetime.now()
55 # Create document record
56 document = Document(
57 filename=data.get("document", json_path.stem),
58 extraction_date=extraction_date,
59 total_pages=data.get("total_pages"),
60 total_tables=data.get("total_tables"),
61 executive_summary=data.get("executive_summary"),
62 model_provider=data.get("model_provider"),
63 model=data.get("model"),
64 )
66 # Initialize embedding service if needed
67 if generate_embeddings and embedding_service is None:
68 logger.info("Initializing embedding service...")
69 embedding_service = get_embedding_service()
71 # Create table records
72 tables_data = data.get("tables", [])
73 for table_data in tables_data:
74 llm_title = table_data.get("llm_title", "")
75 llm_description = table_data.get("llm_description", "")
77 # Generate embeddings if requested
78 title_emb = None
79 desc_emb = None
80 if generate_embeddings and llm_title and llm_description:
81 assert embedding_service is not None
82 # Combined embedding for title + description
83 combined_emb = embedding_service.embed_table_metadata(llm_title, llm_description)
84 # Store same embedding for both fields (can refine later)
85 title_emb = combined_emb.tolist()
86 desc_emb = combined_emb.tolist()
88 table = ExtractedTable(
89 table_id=table_data.get("table_id"),
90 page_number=table_data.get("page"),
91 procedural_title=table_data.get("procedural_title"),
92 llm_title=llm_title,
93 llm_description=llm_description,
94 table_markdown=table_data.get("table_markdown"),
95 headers={"headers": table_data.get("headers", [])},
96 table_metadata=table_data.get("metadata", {}),
97 title_embedding=title_emb,
98 description_embedding=desc_emb,
99 )
100 document.tables.append(table)
102 # Track table count before session operations
103 table_count = len(tables_data)
105 # Save to database
106 if session is None:
107 # Use context manager when no session provided
108 with get_session() as db_session:
109 # Check if document already exists
110 existing = db_session.query(Document).filter_by(filename=document.filename).first()
111 if existing:
112 logger.warning(f"Document {document.filename} already exists (id={existing.id}), skipping")
113 existing_count = db_session.query(ExtractedTable).filter_by(document_id=existing.id).count()
114 return (existing, existing_count)
116 db_session.add(document)
117 db_session.commit()
118 db_session.refresh(document)
120 logger.success(
121 f"Imported {document.filename}: {table_count} tables from {document.total_pages} pages"
122 )
123 return (document, table_count)
124 else:
125 # Use provided session
126 try:
127 existing = session.query(Document).filter_by(filename=document.filename).first()
128 if existing:
129 logger.warning(f"Document {document.filename} already exists (id={existing.id}), skipping")
130 existing_count = session.query(ExtractedTable).filter_by(document_id=existing.id).count()
131 return (existing, existing_count)
133 session.add(document)
134 session.commit()
135 session.refresh(document)
137 logger.success(
138 f"Imported {document.filename}: {table_count} tables from {document.total_pages} pages"
139 )
140 return (document, table_count)
142 except Exception as e:
143 session.rollback()
144 logger.error(f"Failed to import {json_path.name}: {e}")
145 raise
148def import_directory(
149 directory: Path | str,
150 pattern: str = "*_analysis.json",
151 session: Session | None = None,
152 generate_embeddings: bool = False,
153) -> tuple[list[Document], int]:
154 """
155 Import all JSON files matching a pattern from a directory.
157 Args:
158 directory: Path to directory containing JSON files.
159 pattern: Glob pattern for matching files (default: "*_analysis.json").
160 session: Optional database session. If None, creates a new session.
161 generate_embeddings: Whether to generate embeddings during import.
163 Returns:
164 Tuple of (list of Document instances, total table count).
165 """
166 directory = Path(directory)
167 if not directory.exists():
168 raise FileNotFoundError(f"Directory not found: {directory}")
170 json_files = list(directory.glob(pattern))
171 if not json_files:
172 logger.warning(f"No files matching '{pattern}' found in {directory}")
173 return ([], 0)
175 logger.info(f"Found {len(json_files)} JSON files to import")
177 # Initialize embedding service once for all files if needed
178 embedding_service = None
179 if generate_embeddings:
180 logger.info("Initializing embedding service for batch import...")
181 embedding_service = get_embedding_service()
183 documents = []
184 total_tables = 0
186 if session is None:
187 # Import each file with its own session
188 for json_file in json_files:
189 try:
190 doc, table_count = import_json_file(
191 json_file,
192 session=None,
193 embedding_service=embedding_service,
194 generate_embeddings=generate_embeddings,
195 )
196 documents.append(doc)
197 total_tables += table_count
198 except Exception as e:
199 logger.error(f"Failed to import {json_file.name}: {e}")
200 continue
201 else:
202 # Use provided session for all imports
203 for json_file in json_files:
204 try:
205 doc, table_count = import_json_file(
206 json_file,
207 session=session,
208 embedding_service=embedding_service,
209 generate_embeddings=generate_embeddings,
210 )
211 documents.append(doc)
212 total_tables += table_count
213 except Exception as e:
214 logger.error(f"Failed to import {json_file.name}: {e}")
215 continue
217 logger.success(f"Successfully imported {len(documents)}/{len(json_files)} files")
218 return (documents, total_tables)
221def get_import_stats(session: Session | None = None) -> dict[str, Any]:
222 """
223 Get statistics about imported data.
225 Args:
226 session: Optional database session. If None, creates a new session.
228 Returns:
229 Dictionary with import statistics.
230 """
231 if session is None:
232 with get_session() as db_session:
233 doc_count = db_session.query(Document).count()
234 table_count = db_session.query(ExtractedTable).count()
236 # Get documents with table counts
237 documents = db_session.query(Document).all()
239 stats = {
240 "total_documents": doc_count,
241 "total_tables": table_count,
242 "avg_tables_per_doc": round(table_count / doc_count, 2) if doc_count > 0 else 0,
243 "documents": [
244 {"filename": doc.filename, "tables": len(doc.tables), "pages": doc.total_pages}
245 for doc in documents
246 ],
247 }
249 return stats
250 else:
251 doc_count = session.query(Document).count()
252 table_count = session.query(ExtractedTable).count()
254 # Get documents with table counts
255 documents = session.query(Document).all()
257 stats = {
258 "total_documents": doc_count,
259 "total_tables": table_count,
260 "avg_tables_per_doc": round(table_count / doc_count, 2) if doc_count > 0 else 0,
261 "documents": [
262 {"filename": doc.filename, "tables": len(doc.tables), "pages": doc.total_pages}
263 for doc in documents
264 ],
265 }
267 return stats