codekingpro/portable-devtools
114k
1from typing import Any, Callable, Dict, Iterator, List, Optional, Sequence, Union2 3from sqlalchemy.engine import RowMapping4from sqlalchemy.sql.expression import Select5 6from langchain_community.docstore.document import Document7from langchain_community.document_loaders.base import BaseLoader8from langchain_community.utilities.sql_database import SQLDatabase9 10 11class SQLDatabaseLoader(BaseLoader):12 """13 Load documents by querying database tables supported by SQLAlchemy.14 15 For talking to the database, the document loader uses the `SQLDatabase`16 utility from the LangChain integration toolkit.17 18 Each document represents one row of the result.19 """20 21 def __init__(22 self,23 query: Union[str, Select],24 db: SQLDatabase,25 *,26 parameters: Optional[Dict[str, Any]] = None,27 page_content_mapper: Optional[Callable[..., str]] = None,28 metadata_mapper: Optional[Callable[..., Dict[str, Any]]] = None,29 source_columns: Optional[Sequence[str]] = None,30 include_rownum_into_metadata: bool = False,31 include_query_into_metadata: bool = False,32 ):33 """34 Args:35 query: The query to execute.36 db: A LangChain `SQLDatabase`, wrapping an SQLAlchemy engine.37 sqlalchemy_kwargs: More keyword arguments for SQLAlchemy's `create_engine`.38 parameters: Optional. Parameters to pass to the query.39 page_content_mapper: Optional. Function to convert a row into a string40 to use as the `page_content` of the document. By default, the loader41 serializes the whole row into a string, including all columns.42 metadata_mapper: Optional. Function to convert a row into a dictionary43 to use as the `metadata` of the document. By default, no columns are44 selected into the metadata dictionary.45 source_columns: Optional. The names of the columns to use as the `source`46 within the metadata dictionary.47 include_rownum_into_metadata: Optional. Whether to include the row number48 into the metadata dictionary. Default: False.49 include_query_into_metadata: Optional. Whether to include the query50 expression into the metadata dictionary. Default: False.51 """52 self.query = query53 self.db: SQLDatabase = db54 self.parameters = parameters or {}55 self.page_content_mapper = (56 page_content_mapper or self.page_content_default_mapper57 )58 self.metadata_mapper = metadata_mapper or self.metadata_default_mapper59 self.source_columns = source_columns60 self.include_rownum_into_metadata = include_rownum_into_metadata61 self.include_query_into_metadata = include_query_into_metadata62 63 def lazy_load(self) -> Iterator[Document]:64 try:65 import sqlalchemy as sa66 except ImportError:67 raise ImportError(68 "Could not import sqlalchemy python package. "69 "Please install it with `pip install sqlalchemy`."70 )71 72 # Querying in `cursor` fetch mode will return an SQLAlchemy `Result` instance.73 result: sa.Result[Any]74 75 # Invoke the database query.76 if isinstance(self.query, sa.SelectBase):77 result = self.db._execute( # type: ignore[assignment]78 self.query, fetch="cursor", parameters=self.parameters79 )80 query_sql = str(self.query.compile(bind=self.db._engine))81 elif isinstance(self.query, str):82 result = self.db._execute( # type: ignore[assignment]83 sa.text(self.query), fetch="cursor", parameters=self.parameters84 )85 query_sql = self.query86 else:87 raise TypeError(f"Unable to process query of unknown type: {self.query}")88 89 # Iterate database result rows and generate list of documents.90 for i, row in enumerate(result.mappings()):91 page_content = self.page_content_mapper(row)92 metadata = self.metadata_mapper(row)93 94 if self.include_rownum_into_metadata:95 metadata["row"] = i96 if self.include_query_into_metadata:97 metadata["query"] = query_sql98 99 source_values = []100 for column, value in row.items():101 if self.source_columns and column in self.source_columns:102 source_values.append(value)103 if source_values:104 metadata["source"] = ",".join(source_values)105 106 yield Document(page_content=page_content, metadata=metadata)107 108 @staticmethod109 def page_content_default_mapper(110 row: RowMapping, column_names: Optional[List[str]] = None111 ) -> str:112 """113 A reasonable default function to convert a record into a "page content" string.114 """115 if column_names is None:116 column_names = list(row.keys())117 return "\n".join(118 f"{column}: {value}"119 for column, value in row.items()120 if column in column_names121 )122 123 @staticmethod124 def metadata_default_mapper(125 row: RowMapping, column_names: Optional[List[str]] = None126 ) -> Dict[str, Any]:127 """128 A reasonable default function to convert a record into a "metadata" dictionary.129 """130 if column_names is None:131 return {}132 133 metadata: Dict[str, Any] = {}134 for column, value in row.items():135 if column in column_names:136 metadata[column] = value137 return metadata138 