codekingpro/portable-devtools
114k
1import asyncio2import logging3from typing import AsyncIterator, Dict, Iterator, List, Optional, Sequence4 5from langchain_core.documents import Document6 7from langchain_community.document_loaders.base import BaseLoader8 9logger = logging.getLogger(__name__)10 11 12class MongodbLoader(BaseLoader):13 """Load MongoDB documents."""14 15 def __init__(16 self,17 connection_string: str,18 db_name: str,19 collection_name: str,20 *,21 filter_criteria: Optional[Dict] = None,22 field_names: Optional[Sequence[str]] = None,23 metadata_names: Optional[Sequence[str]] = None,24 include_db_collection_in_metadata: bool = True,25 ) -> None:26 """27 Initializes the MongoDB loader with necessary database connection28 details and configurations.29 30 Args:31 connection_string (str): MongoDB connection URI.32 db_name (str):Name of the database to connect to.33 collection_name (str): Name of the collection to fetch documents from.34 filter_criteria (Optional[Dict]): MongoDB filter criteria for querying35 documents.36 field_names (Optional[Sequence[str]]): List of field names to retrieve37 from documents.38 metadata_names (Optional[Sequence[str]]): Additional metadata fields to39 extract from documents.40 include_db_collection_in_metadata (bool): Flag to include database and41 collection names in metadata.42 43 Raises:44 ImportError: If the motor library is not installed.45 ValueError: If any necessary argument is missing.46 """47 try:48 from motor.motor_asyncio import AsyncIOMotorClient49 except ImportError as e:50 raise ImportError(51 "Cannot import from motor, please install with `pip install motor`."52 ) from e53 54 if not connection_string:55 raise ValueError("connection_string must be provided.")56 57 if not db_name:58 raise ValueError("db_name must be provided.")59 60 if not collection_name:61 raise ValueError("collection_name must be provided.")62 63 self.client = AsyncIOMotorClient(connection_string)64 self.db_name = db_name65 self.collection_name = collection_name66 self.field_names = field_names or []67 self.filter_criteria = filter_criteria or {}68 self.metadata_names = metadata_names or []69 self.include_db_collection_in_metadata = include_db_collection_in_metadata70 71 self.db = self.client.get_database(db_name)72 self.collection = self.db.get_collection(collection_name)73 74 def load(self) -> List[Document]:75 """Load data into Document objects.76 77 Attention:78 79 This implementation starts an asyncio event loop which80 will only work if running in a sync env. In an async env, it should81 fail since there is already an event loop running.82 83 This code should be updated to kick off the event loop from a separate84 thread if running within an async context.85 """86 return asyncio.run(self.aload())87 88 def lazy_load(self) -> Iterator[Document]:89 """A lazy loader for MongoDB documents.90 91 Attention:92 93 This implementation starts an asyncio event loop which94 will only work if running in a sync env. In an async env, it should95 fail since there is already an event loop running.96 97 This code should be updated to kick off the event loop from a separate98 thread if running within an async context.99 100 Yields:101 Document: A document from the MongoDB collection.102 """103 try:104 event_loop = asyncio.get_running_loop()105 except RuntimeError:106 event_loop = asyncio.new_event_loop()107 asyncio.set_event_loop(event_loop)108 109 async_generator = self.alazy_load()110 111 while True:112 try:113 document = event_loop.run_until_complete(async_generator.__anext__())114 yield document115 except StopAsyncIteration:116 break117 118 async def alazy_load(self) -> AsyncIterator[Document]:119 """Asynchronously yields Document objects one at a time.120 121 Yields:122 Document: A document from the MongoDB collection.123 """124 projection = self._construct_projection()125 126 async for doc in self.collection.find(self.filter_criteria, projection):127 yield self._process_document(doc)128 129 async def aload(self) -> List[Document]:130 """Asynchronously loads data into Document objects."""131 result = []132 total_docs = await self.collection.count_documents(self.filter_criteria)133 134 projection = self._construct_projection()135 136 async for doc in self.collection.find(self.filter_criteria, projection):137 result.append(self._process_document(doc))138 139 if len(result) != total_docs:140 logger.warning(141 f"Only partial collection of documents returned. "142 f"Loaded {len(result)} docs, expected {total_docs}."143 )144 145 return result146 147 def _process_document(self, doc: Dict) -> Document:148 """Process a single MongoDB document into a Document object.149 150 Args:151 doc: The MongoDB document dictionary to process into a Document object.152 """153 metadata = self._extract_fields(doc, self.metadata_names, default="")154 155 # Optionally add database and collection names to metadata156 if self.include_db_collection_in_metadata:157 metadata.update(158 {159 "database": self.db_name,160 "collection": self.collection_name,161 }162 )163 164 # Extract text content from filtered fields or use the entire document165 if self.field_names is not None:166 fields = self._extract_fields(doc, self.field_names, default="")167 texts = [str(value) for value in fields.values()]168 text = " ".join(texts)169 else:170 text = str(doc)171 172 return Document(page_content=text, metadata=metadata)173 174 def _construct_projection(self) -> Optional[Dict]:175 """Constructs the projection dictionary for MongoDB query based176 on the specified field names and metadata names."""177 field_names = list(self.field_names) or []178 metadata_names = list(self.metadata_names) or []179 all_fields = field_names + metadata_names180 return {field: 1 for field in all_fields} if all_fields else None181 182 def _extract_fields(183 self,184 document: Dict,185 fields: Sequence[str],186 default: str = "",187 ) -> Dict:188 """Extracts and returns values for specified fields from a document."""189 extracted = {}190 for field in fields or []:191 value = document192 for key in field.split("."):193 value = value.get(key, default)194 if value == default:195 break196 new_field_name = field.replace(".", "_")197 extracted[new_field_name] = value198 return extracted199 