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

1"""Import utility for loading JSON analyses into the database.""" 

2 

3import json 

4from datetime import datetime 

5from pathlib import Path 

6from typing import Any 

7 

8from loguru import logger 

9from sqlalchemy.orm import Session 

10 

11from quber.db.connection import get_session 

12from quber.db.embeddings import EmbeddingService, get_embedding_service 

13from quber.db.models import Document, ExtractedTable 

14 

15 

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. 

24 

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. 

30 

31 Returns: 

32 Tuple of (Document instance, number of tables imported). 

33 

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}") 

41 

42 logger.info(f"Importing {json_path.name}...") 

43 

44 with open(json_path) as f: 

45 data: dict[str, Any] = json.load(f) 

46 

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

54 

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 ) 

65 

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

70 

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", "") 

76 

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

87 

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) 

101 

102 # Track table count before session operations 

103 table_count = len(tables_data) 

104 

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) 

115 

116 db_session.add(document) 

117 db_session.commit() 

118 db_session.refresh(document) 

119 

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) 

132 

133 session.add(document) 

134 session.commit() 

135 session.refresh(document) 

136 

137 logger.success( 

138 f"Imported {document.filename}: {table_count} tables from {document.total_pages} pages" 

139 ) 

140 return (document, table_count) 

141 

142 except Exception as e: 

143 session.rollback() 

144 logger.error(f"Failed to import {json_path.name}: {e}") 

145 raise 

146 

147 

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. 

156 

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. 

162 

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}") 

169 

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) 

174 

175 logger.info(f"Found {len(json_files)} JSON files to import") 

176 

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

182 

183 documents = [] 

184 total_tables = 0 

185 

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 

216 

217 logger.success(f"Successfully imported {len(documents)}/{len(json_files)} files") 

218 return (documents, total_tables) 

219 

220 

221def get_import_stats(session: Session | None = None) -> dict[str, Any]: 

222 """ 

223 Get statistics about imported data. 

224 

225 Args: 

226 session: Optional database session. If None, creates a new session. 

227 

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

235 

236 # Get documents with table counts 

237 documents = db_session.query(Document).all() 

238 

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 } 

248 

249 return stats 

250 else: 

251 doc_count = session.query(Document).count() 

252 table_count = session.query(ExtractedTable).count() 

253 

254 # Get documents with table counts 

255 documents = session.query(Document).all() 

256 

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 } 

266 

267 return stats