management_service.py 38 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643644645646647648649650651652653654655656657658659660661662663664665666667668669670671672673674675676677678679680681682683684685686687688689690691692693694695696697698699700701702703704705706707708709710711712713714715716717718719720721722723724725726727728729730731732733734735736737738739740741742743744745746747748749750751752753754755756757758759760761762763764765766767768769770771772773774775776777778779780781782783784785786787788789790791792793794795796797798799800801802803804805806807808809810811812813814815816817818819820821822823824825826827828829830831832833834835836837838839840841842843844845846847848849850851852853854855856857858859860861862863864865866867868869870871872873874875876877878879880881882883884885886887888889890891892893894895896897898899900901902903904905906907908909910911912913914915916917918919920921922923924925926927928929930931932933934935936937938939940941942943944945946947948949950951952953954955956957958959960961962963964965966967968969970971972973974975976977978979980981982983984985986987988989990991992993994995996997998999100010011002100310041005100610071008100910101011101210131014101510161017101810191020102110221023102410251026102710281029103010311032103310341035103610371038103910401041104210431044104510461047104810491050105110521053105410551056105710581059106010611062106310641065106610671068106910701071107210731074107510761077107810791080108110821083108410851086108710881089109010911092109310941095109610971098
  1. import logging
  2. import os
  3. from collections import defaultdict
  4. from datetime import datetime, timedelta, timezone
  5. from typing import IO, Any, BinaryIO, Optional, Tuple
  6. from uuid import UUID
  7. import toml
  8. from core.base import (
  9. CollectionResponse,
  10. ConversationResponse,
  11. DocumentResponse,
  12. GenerationConfig,
  13. GraphConstructionStatus,
  14. Message,
  15. MessageResponse,
  16. Prompt,
  17. R2RException,
  18. StoreType,
  19. User,
  20. )
  21. from ..abstractions import R2RProviders
  22. from ..config import R2RConfig
  23. from .base import Service
  24. logger = logging.getLogger()
  25. class ManagementService(Service):
  26. def __init__(
  27. self,
  28. config: R2RConfig,
  29. providers: R2RProviders,
  30. ):
  31. super().__init__(
  32. config,
  33. providers,
  34. )
  35. async def app_settings(self):
  36. prompts = (
  37. await self.providers.database.prompts_handler.get_all_prompts()
  38. )
  39. config_toml = self.config.to_toml()
  40. config_dict = toml.loads(config_toml)
  41. try:
  42. project_name = os.environ["R2R_PROJECT_NAME"]
  43. except KeyError:
  44. project_name = ""
  45. return {
  46. "config": config_dict,
  47. "prompts": prompts,
  48. "r2r_project_name": project_name,
  49. }
  50. async def users_overview(
  51. self,
  52. offset: int,
  53. limit: int,
  54. user_ids: Optional[list[UUID]] = None,
  55. ):
  56. return await self.providers.database.users_handler.get_users_overview(
  57. offset=offset,
  58. limit=limit,
  59. user_ids=user_ids,
  60. )
  61. async def delete_documents_and_chunks_by_filter(
  62. self,
  63. filters: dict[str, Any],
  64. ):
  65. """Delete chunks matching the given filters. If any documents are now
  66. empty (i.e., have no remaining chunks), delete those documents as well.
  67. Args:
  68. filters (dict[str, Any]): Filters specifying which chunks to delete.
  69. chunks_handler (PostgresChunksHandler): The handler for chunk operations.
  70. documents_handler (PostgresDocumentsHandler): The handler for document operations.
  71. graphs_handler: Handler for entity and relationship operations in the Graph.
  72. Returns:
  73. dict: A summary of what was deleted.
  74. """
  75. def transform_chunk_id_to_id(
  76. filters: dict[str, Any],
  77. ) -> dict[str, Any]:
  78. """Example transformation function if your filters use `chunk_id`
  79. instead of `id`.
  80. Recursively transform `chunk_id` to `id`.
  81. """
  82. if isinstance(filters, dict):
  83. transformed = {}
  84. for key, value in filters.items():
  85. if key == "chunk_id":
  86. transformed["id"] = value
  87. elif key in ["$and", "$or"]:
  88. transformed[key] = [
  89. transform_chunk_id_to_id(item) for item in value
  90. ]
  91. else:
  92. transformed[key] = transform_chunk_id_to_id(value)
  93. return transformed
  94. return filters
  95. # Transform filters if needed.
  96. transformed_filters = transform_chunk_id_to_id(filters)
  97. # Find chunks that match the filters before deleting
  98. interim_results = (
  99. await self.providers.database.chunks_handler.list_chunks(
  100. filters=transformed_filters,
  101. offset=0,
  102. limit=1_000,
  103. include_vectors=False,
  104. )
  105. )
  106. results = interim_results["results"]
  107. while interim_results["total_entries"] == 1_000:
  108. # If we hit the limit, we need to paginate to get all results
  109. interim_results = (
  110. await self.providers.database.chunks_handler.list_chunks(
  111. filters=transformed_filters,
  112. offset=interim_results["offset"] + 1_000,
  113. limit=1_000,
  114. include_vectors=False,
  115. )
  116. )
  117. results.extend(interim_results["results"])
  118. document_ids = set()
  119. owner_id = None
  120. if "$and" in filters:
  121. for condition in filters["$and"]:
  122. if "owner_id" in condition and "$eq" in condition["owner_id"]:
  123. owner_id = condition["owner_id"]["$eq"]
  124. elif (
  125. "document_id" in condition
  126. and "$eq" in condition["document_id"]
  127. ):
  128. document_ids.add(UUID(condition["document_id"]["$eq"]))
  129. elif "document_id" in filters:
  130. doc_id = filters["document_id"]
  131. if isinstance(doc_id, str):
  132. document_ids.add(UUID(doc_id))
  133. elif isinstance(doc_id, UUID):
  134. document_ids.add(doc_id)
  135. elif isinstance(doc_id, dict) and "$eq" in doc_id:
  136. value = doc_id["$eq"]
  137. document_ids.add(
  138. UUID(value) if isinstance(value, str) else value
  139. )
  140. # Delete matching chunks from the database
  141. delete_results = await self.providers.database.chunks_handler.delete(
  142. transformed_filters
  143. )
  144. # Extract the document_ids that were affected.
  145. affected_doc_ids = {
  146. UUID(info["document_id"])
  147. for info in delete_results.values()
  148. if info.get("document_id")
  149. }
  150. document_ids.update(affected_doc_ids)
  151. # Check if the document still has any chunks left
  152. docs_to_delete = []
  153. for doc_id in document_ids:
  154. documents_overview_response = await self.providers.database.documents_handler.get_documents_overview(
  155. offset=0, limit=1, filter_document_ids=[doc_id]
  156. )
  157. if not documents_overview_response["results"]:
  158. raise R2RException(
  159. status_code=404, message="Document not found"
  160. )
  161. document = documents_overview_response["results"][0]
  162. for collection_id in document.collection_ids:
  163. await self.providers.database.collections_handler.decrement_collection_document_count(
  164. collection_id=collection_id
  165. )
  166. if owner_id and str(document.owner_id) != owner_id:
  167. raise R2RException(
  168. status_code=404,
  169. message="Document not found or insufficient permissions",
  170. )
  171. docs_to_delete.append(doc_id)
  172. # Delete documents that no longer have associated chunks
  173. for doc_id in docs_to_delete:
  174. # Delete related entities & relationships if needed:
  175. await self.providers.database.graphs_handler.entities.delete(
  176. parent_id=doc_id,
  177. store_type=StoreType.DOCUMENTS,
  178. )
  179. await self.providers.database.graphs_handler.relationships.delete(
  180. parent_id=doc_id,
  181. store_type=StoreType.DOCUMENTS,
  182. )
  183. # Finally, delete the document from documents_overview:
  184. await self.providers.database.documents_handler.delete(
  185. document_id=doc_id
  186. )
  187. return {
  188. "success": True,
  189. "deleted_chunks_count": len(delete_results),
  190. "deleted_documents_count": len(docs_to_delete),
  191. "deleted_document_ids": [str(d) for d in docs_to_delete],
  192. }
  193. async def download_file(
  194. self, document_id: UUID
  195. ) -> Optional[Tuple[str, BinaryIO, int]]:
  196. if result := await self.providers.file.retrieve_file(document_id):
  197. return result
  198. return None
  199. async def export_files(
  200. self,
  201. document_ids: Optional[list[UUID]] = None,
  202. start_date: Optional[datetime] = None,
  203. end_date: Optional[datetime] = None,
  204. ) -> tuple[str, BinaryIO, int]:
  205. return await self.providers.file.retrieve_files_as_zip(
  206. document_ids=document_ids,
  207. start_date=start_date,
  208. end_date=end_date,
  209. )
  210. async def export_collections(
  211. self,
  212. columns: Optional[list[str]] = None,
  213. filters: Optional[dict] = None,
  214. include_header: bool = True,
  215. ) -> tuple[str, IO]:
  216. return await self.providers.database.collections_handler.export_to_csv(
  217. columns=columns,
  218. filters=filters,
  219. include_header=include_header,
  220. )
  221. async def export_documents(
  222. self,
  223. columns: Optional[list[str]] = None,
  224. filters: Optional[dict] = None,
  225. include_header: bool = True,
  226. ) -> tuple[str, IO]:
  227. return await self.providers.database.documents_handler.export_to_csv(
  228. columns=columns,
  229. filters=filters,
  230. include_header=include_header,
  231. )
  232. async def export_document_entities(
  233. self,
  234. id: UUID,
  235. columns: Optional[list[str]] = None,
  236. filters: Optional[dict] = None,
  237. include_header: bool = True,
  238. ) -> tuple[str, IO]:
  239. return await self.providers.database.graphs_handler.entities.export_to_csv(
  240. parent_id=id,
  241. store_type=StoreType.DOCUMENTS,
  242. columns=columns,
  243. filters=filters,
  244. include_header=include_header,
  245. )
  246. async def export_document_relationships(
  247. self,
  248. id: UUID,
  249. columns: Optional[list[str]] = None,
  250. filters: Optional[dict] = None,
  251. include_header: bool = True,
  252. ) -> tuple[str, IO]:
  253. return await self.providers.database.graphs_handler.relationships.export_to_csv(
  254. parent_id=id,
  255. store_type=StoreType.DOCUMENTS,
  256. columns=columns,
  257. filters=filters,
  258. include_header=include_header,
  259. )
  260. async def export_conversations(
  261. self,
  262. columns: Optional[list[str]] = None,
  263. filters: Optional[dict] = None,
  264. include_header: bool = True,
  265. ) -> tuple[str, IO]:
  266. return await self.providers.database.conversations_handler.export_conversations_to_csv(
  267. columns=columns,
  268. filters=filters,
  269. include_header=include_header,
  270. )
  271. async def export_graph_entities(
  272. self,
  273. id: UUID,
  274. columns: Optional[list[str]] = None,
  275. filters: Optional[dict] = None,
  276. include_header: bool = True,
  277. ) -> tuple[str, IO]:
  278. return await self.providers.database.graphs_handler.entities.export_to_csv(
  279. parent_id=id,
  280. store_type=StoreType.GRAPHS,
  281. columns=columns,
  282. filters=filters,
  283. include_header=include_header,
  284. )
  285. async def export_graph_relationships(
  286. self,
  287. id: UUID,
  288. columns: Optional[list[str]] = None,
  289. filters: Optional[dict] = None,
  290. include_header: bool = True,
  291. ) -> tuple[str, IO]:
  292. return await self.providers.database.graphs_handler.relationships.export_to_csv(
  293. parent_id=id,
  294. store_type=StoreType.GRAPHS,
  295. columns=columns,
  296. filters=filters,
  297. include_header=include_header,
  298. )
  299. async def export_graph_communities(
  300. self,
  301. id: UUID,
  302. columns: Optional[list[str]] = None,
  303. filters: Optional[dict] = None,
  304. include_header: bool = True,
  305. ) -> tuple[str, IO]:
  306. return await self.providers.database.graphs_handler.communities.export_to_csv(
  307. parent_id=id,
  308. store_type=StoreType.GRAPHS,
  309. columns=columns,
  310. filters=filters,
  311. include_header=include_header,
  312. )
  313. async def export_messages(
  314. self,
  315. columns: Optional[list[str]] = None,
  316. filters: Optional[dict] = None,
  317. include_header: bool = True,
  318. ) -> tuple[str, IO]:
  319. return await self.providers.database.conversations_handler.export_messages_to_csv(
  320. columns=columns,
  321. filters=filters,
  322. include_header=include_header,
  323. )
  324. async def export_users(
  325. self,
  326. columns: Optional[list[str]] = None,
  327. filters: Optional[dict] = None,
  328. include_header: bool = True,
  329. ) -> tuple[str, IO]:
  330. return await self.providers.database.users_handler.export_to_csv(
  331. columns=columns,
  332. filters=filters,
  333. include_header=include_header,
  334. )
  335. async def documents_overview(
  336. self,
  337. offset: int,
  338. limit: int,
  339. user_ids: Optional[list[UUID]] = None,
  340. collection_ids: Optional[list[UUID]] = None,
  341. document_ids: Optional[list[UUID]] = None,
  342. owner_only: bool = False,
  343. ):
  344. return await self.providers.database.documents_handler.get_documents_overview(
  345. offset=offset,
  346. limit=limit,
  347. filter_document_ids=document_ids,
  348. filter_user_ids=user_ids,
  349. filter_collection_ids=collection_ids,
  350. owner_only=owner_only,
  351. )
  352. async def update_document_metadata(
  353. self,
  354. document_id: UUID,
  355. metadata: list[dict],
  356. overwrite: bool = False,
  357. ):
  358. return await self.providers.database.documents_handler.update_document_metadata(
  359. document_id=document_id,
  360. metadata=metadata,
  361. overwrite=overwrite,
  362. )
  363. async def list_document_chunks(
  364. self,
  365. document_id: UUID,
  366. offset: int,
  367. limit: int,
  368. include_vectors: bool = False,
  369. ):
  370. return (
  371. await self.providers.database.chunks_handler.list_document_chunks(
  372. document_id=document_id,
  373. offset=offset,
  374. limit=limit,
  375. include_vectors=include_vectors,
  376. )
  377. )
  378. async def assign_document_to_collection(
  379. self, document_id: UUID, collection_id: UUID
  380. ):
  381. await self.providers.database.chunks_handler.assign_document_chunks_to_collection(
  382. document_id, collection_id
  383. )
  384. await self.providers.database.collections_handler.assign_document_to_collection_relational(
  385. document_id, collection_id
  386. )
  387. await self.providers.database.documents_handler.set_workflow_status(
  388. id=collection_id,
  389. status_type="graph_sync_status",
  390. status=GraphConstructionStatus.OUTDATED,
  391. )
  392. await self.providers.database.documents_handler.set_workflow_status(
  393. id=collection_id,
  394. status_type="graph_cluster_status",
  395. status=GraphConstructionStatus.OUTDATED,
  396. )
  397. return {"message": "Document assigned to collection successfully"}
  398. async def remove_document_from_collection(
  399. self, document_id: UUID, collection_id: UUID
  400. ):
  401. await self.providers.database.collections_handler.remove_document_from_collection_relational(
  402. document_id, collection_id
  403. )
  404. await self.providers.database.chunks_handler.remove_document_from_collection_vector(
  405. document_id, collection_id
  406. )
  407. # await self.providers.database.graphs_handler.delete_node_via_document_id(
  408. # document_id, collection_id
  409. # )
  410. return None
  411. def _process_relationships(
  412. self, relationships: list[Tuple[str, str, str]]
  413. ) -> Tuple[dict[str, list[str]], dict[str, dict[str, list[str]]]]:
  414. graph = defaultdict(list)
  415. grouped: dict[str, dict[str, list[str]]] = defaultdict(
  416. lambda: defaultdict(list)
  417. )
  418. for subject, relation, obj in relationships:
  419. graph[subject].append(obj)
  420. grouped[subject][relation].append(obj)
  421. if obj not in graph:
  422. graph[obj] = []
  423. return dict(graph), dict(grouped)
  424. def generate_output(
  425. self,
  426. grouped_relationships: dict[str, dict[str, list[str]]],
  427. graph: dict[str, list[str]],
  428. descriptions_dict: dict[str, str],
  429. print_descriptions: bool = True,
  430. ) -> list[str]:
  431. output = []
  432. # Print grouped relationships
  433. for subject, relations in grouped_relationships.items():
  434. output.append(f"\n== {subject} ==")
  435. if print_descriptions and subject in descriptions_dict:
  436. output.append(f"\tDescription: {descriptions_dict[subject]}")
  437. for relation, objects in relations.items():
  438. output.append(f" {relation}:")
  439. for obj in objects:
  440. output.append(f" - {obj}")
  441. if print_descriptions and obj in descriptions_dict:
  442. output.append(
  443. f" Description: {descriptions_dict[obj]}"
  444. )
  445. # Print basic graph statistics
  446. output.extend(
  447. [
  448. "\n== Graph Statistics ==",
  449. f"Number of nodes: {len(graph)}",
  450. f"Number of edges: {sum(len(neighbors) for neighbors in graph.values())}",
  451. f"Number of connected components: {self._count_connected_components(graph)}",
  452. ]
  453. )
  454. # Find central nodes
  455. central_nodes = self._get_central_nodes(graph)
  456. output.extend(
  457. [
  458. "\n== Most Central Nodes ==",
  459. *(
  460. f" {node}: {centrality:.4f}"
  461. for node, centrality in central_nodes
  462. ),
  463. ]
  464. )
  465. return output
  466. def _count_connected_components(self, graph: dict[str, list[str]]) -> int:
  467. visited = set()
  468. components = 0
  469. def dfs(node):
  470. visited.add(node)
  471. for neighbor in graph[node]:
  472. if neighbor not in visited:
  473. dfs(neighbor)
  474. for node in graph:
  475. if node not in visited:
  476. dfs(node)
  477. components += 1
  478. return components
  479. def _get_central_nodes(
  480. self, graph: dict[str, list[str]]
  481. ) -> list[Tuple[str, float]]:
  482. degree = {node: len(neighbors) for node, neighbors in graph.items()}
  483. total_nodes = len(graph)
  484. centrality = {
  485. node: deg / (total_nodes - 1) for node, deg in degree.items()
  486. }
  487. return sorted(centrality.items(), key=lambda x: x[1], reverse=True)[:5]
  488. async def create_collection(
  489. self,
  490. owner_id: UUID,
  491. name: Optional[str] = None,
  492. description: str | None = None,
  493. ) -> CollectionResponse:
  494. result = await self.providers.database.collections_handler.create_collection(
  495. owner_id=owner_id,
  496. name=name,
  497. description=description,
  498. )
  499. await self.providers.database.graphs_handler.create(
  500. collection_id=result.id,
  501. name=name,
  502. description=description,
  503. )
  504. return result
  505. async def update_collection(
  506. self,
  507. collection_id: UUID,
  508. name: Optional[str] = None,
  509. description: Optional[str] = None,
  510. generate_description: bool = False,
  511. ) -> CollectionResponse:
  512. if generate_description:
  513. description = await self.summarize_collection(
  514. id=collection_id, offset=0, limit=100
  515. )
  516. return await self.providers.database.collections_handler.update_collection(
  517. collection_id=collection_id,
  518. name=name,
  519. description=description,
  520. )
  521. async def delete_collection(self, collection_id: UUID) -> bool:
  522. logger.info(f"Deleting collection {collection_id}")
  523. await self.providers.database.collections_handler.delete_collection_relational(
  524. collection_id
  525. )
  526. logger.info(f"Deleting collection {collection_id} from chunks")
  527. await self.providers.database.chunks_handler.delete_collection_vector(
  528. collection_id
  529. )
  530. try:
  531. logger.info(f"Deleting collection {collection_id} from graph")
  532. await self.providers.database.graphs_handler.delete(
  533. collection_id=collection_id,
  534. )
  535. except Exception as e:
  536. logger.warning(
  537. f"Error deleting graph for collection {collection_id}: {e}"
  538. )
  539. return True
  540. async def collections_overview(
  541. self,
  542. offset: int,
  543. limit: int,
  544. user_ids: Optional[list[UUID]] = None,
  545. document_ids: Optional[list[UUID]] = None,
  546. collection_ids: Optional[list[UUID]] = None,
  547. owner_only: bool = False,
  548. ) -> dict[str, list[CollectionResponse] | int]:
  549. return await self.providers.database.collections_handler.get_collections_overview(
  550. offset=offset,
  551. limit=limit,
  552. filter_user_ids=user_ids,
  553. filter_document_ids=document_ids,
  554. filter_collection_ids=collection_ids,
  555. owner_only=owner_only,
  556. )
  557. async def add_user_to_collection(
  558. self, user_id: UUID, collection_id: UUID
  559. ) -> bool:
  560. return (
  561. await self.providers.database.users_handler.add_user_to_collection(
  562. user_id, collection_id
  563. )
  564. )
  565. async def remove_user_from_collection(
  566. self, user_id: UUID, collection_id: UUID
  567. ) -> bool:
  568. return await self.providers.database.users_handler.remove_user_from_collection(
  569. user_id, collection_id
  570. )
  571. async def get_users_in_collection(
  572. self, collection_id: UUID, offset: int = 0, limit: int = 100
  573. ) -> dict[str, list[User] | int]:
  574. return await self.providers.database.users_handler.get_users_in_collection(
  575. collection_id, offset=offset, limit=limit
  576. )
  577. async def documents_in_collection(
  578. self, collection_id: UUID, offset: int = 0, limit: int = 100
  579. ) -> dict[str, list[DocumentResponse] | int]:
  580. return await self.providers.database.collections_handler.documents_in_collection(
  581. collection_id, offset=offset, limit=limit
  582. )
  583. async def summarize_collection(
  584. self, id: UUID, offset: int, limit: int
  585. ) -> str:
  586. documents_in_collection_response = await self.documents_in_collection(
  587. collection_id=id,
  588. offset=offset,
  589. limit=limit,
  590. )
  591. document_summaries = [
  592. document.summary
  593. for document in documents_in_collection_response["results"] # type: ignore
  594. ]
  595. logger.info(
  596. f"Summarizing collection {id} with {len(document_summaries)} of {documents_in_collection_response['total_entries']} documents."
  597. )
  598. formatted_summaries = "\n\n".join(document_summaries) # type: ignore
  599. messages = await self.providers.database.prompts_handler.get_message_payload(
  600. system_prompt_name=self.config.database.collection_summary_system_prompt,
  601. task_prompt_name=self.config.database.collection_summary_prompt,
  602. task_inputs={"document_summaries": formatted_summaries},
  603. )
  604. response = await self.providers.llm.aget_completion(
  605. messages=messages,
  606. generation_config=GenerationConfig(
  607. model=self.config.ingestion.document_summary_model
  608. or self.config.app.fast_llm
  609. ),
  610. )
  611. if collection_summary := response.choices[0].message.content:
  612. return collection_summary
  613. else:
  614. raise ValueError("Expected a generated response.")
  615. async def add_prompt(
  616. self, name: str, template: str, input_types: dict[str, str]
  617. ) -> dict:
  618. try:
  619. await self.providers.database.prompts_handler.add_prompt(
  620. name, template, input_types
  621. )
  622. return f"Prompt '{name}' added successfully." # type: ignore
  623. except ValueError as e:
  624. raise R2RException(status_code=400, message=str(e)) from e
  625. async def get_cached_prompt(
  626. self,
  627. prompt_name: str,
  628. inputs: Optional[dict[str, Any]] = None,
  629. prompt_override: Optional[str] = None,
  630. ) -> dict:
  631. try:
  632. return {
  633. "message": (
  634. await self.providers.database.prompts_handler.get_cached_prompt(
  635. prompt_name=prompt_name,
  636. inputs=inputs,
  637. prompt_override=prompt_override,
  638. )
  639. )
  640. }
  641. except ValueError as e:
  642. raise R2RException(status_code=404, message=str(e)) from e
  643. async def get_prompt(
  644. self,
  645. prompt_name: str,
  646. inputs: Optional[dict[str, Any]] = None,
  647. prompt_override: Optional[str] = None,
  648. ) -> dict:
  649. try:
  650. return await self.providers.database.prompts_handler.get_prompt( # type: ignore
  651. name=prompt_name,
  652. inputs=inputs,
  653. prompt_override=prompt_override,
  654. )
  655. except ValueError as e:
  656. raise R2RException(status_code=404, message=str(e)) from e
  657. async def get_all_prompts(self) -> dict[str, Prompt]:
  658. return await self.providers.database.prompts_handler.get_all_prompts()
  659. async def update_prompt(
  660. self,
  661. name: str,
  662. template: Optional[str] = None,
  663. input_types: Optional[dict[str, str]] = None,
  664. ) -> dict:
  665. try:
  666. await self.providers.database.prompts_handler.update_prompt(
  667. name, template, input_types
  668. )
  669. return f"Prompt '{name}' updated successfully." # type: ignore
  670. except ValueError as e:
  671. raise R2RException(status_code=404, message=str(e)) from e
  672. async def delete_prompt(self, name: str) -> dict:
  673. try:
  674. await self.providers.database.prompts_handler.delete_prompt(name)
  675. return {"message": f"Prompt '{name}' deleted successfully."}
  676. except ValueError as e:
  677. raise R2RException(status_code=404, message=str(e)) from e
  678. async def get_conversation(
  679. self,
  680. conversation_id: UUID,
  681. user_ids: Optional[list[UUID]] = None,
  682. ) -> list[MessageResponse]:
  683. return await self.providers.database.conversations_handler.get_conversation(
  684. conversation_id=conversation_id,
  685. filter_user_ids=user_ids,
  686. )
  687. async def create_conversation(
  688. self,
  689. user_id: Optional[UUID] = None,
  690. name: Optional[str] = None,
  691. ) -> ConversationResponse:
  692. return await self.providers.database.conversations_handler.create_conversation(
  693. user_id=user_id,
  694. name=name,
  695. )
  696. async def conversations_overview(
  697. self,
  698. offset: int,
  699. limit: int,
  700. conversation_ids: Optional[list[UUID]] = None,
  701. user_ids: Optional[list[UUID]] = None,
  702. ) -> dict[str, list[dict] | int]:
  703. return await self.providers.database.conversations_handler.get_conversations_overview(
  704. offset=offset,
  705. limit=limit,
  706. filter_user_ids=user_ids,
  707. conversation_ids=conversation_ids,
  708. )
  709. async def add_message(
  710. self,
  711. conversation_id: UUID,
  712. content: Message,
  713. parent_id: Optional[UUID] = None,
  714. metadata: Optional[dict] = None,
  715. ) -> MessageResponse:
  716. return await self.providers.database.conversations_handler.add_message(
  717. conversation_id=conversation_id,
  718. content=content,
  719. parent_id=parent_id,
  720. metadata=metadata,
  721. )
  722. async def edit_message(
  723. self,
  724. message_id: UUID,
  725. new_content: Optional[str] = None,
  726. additional_metadata: Optional[dict] = None,
  727. ) -> dict[str, Any]:
  728. return (
  729. await self.providers.database.conversations_handler.edit_message(
  730. message_id=message_id,
  731. new_content=new_content,
  732. additional_metadata=additional_metadata or {},
  733. )
  734. )
  735. async def update_conversation(
  736. self, conversation_id: UUID, name: str
  737. ) -> ConversationResponse:
  738. return await self.providers.database.conversations_handler.update_conversation(
  739. conversation_id=conversation_id, name=name
  740. )
  741. async def delete_conversation(
  742. self,
  743. conversation_id: UUID,
  744. user_ids: Optional[list[UUID]] = None,
  745. ) -> None:
  746. await (
  747. self.providers.database.conversations_handler.delete_conversation(
  748. conversation_id=conversation_id,
  749. filter_user_ids=user_ids,
  750. )
  751. )
  752. async def get_user_max_documents(self, user_id: UUID) -> int | None:
  753. # Fetch the user to see if they have any overrides stored
  754. user = await self.providers.database.users_handler.get_user_by_id(
  755. user_id
  756. )
  757. if user.limits_overrides and "max_documents" in user.limits_overrides:
  758. return user.limits_overrides["max_documents"]
  759. return self.config.app.default_max_documents_per_user
  760. async def get_user_max_chunks(self, user_id: UUID) -> int | None:
  761. user = await self.providers.database.users_handler.get_user_by_id(
  762. user_id
  763. )
  764. if user.limits_overrides and "max_chunks" in user.limits_overrides:
  765. return user.limits_overrides["max_chunks"]
  766. return self.config.app.default_max_chunks_per_user
  767. async def get_user_max_collections(self, user_id: UUID) -> int | None:
  768. user = await self.providers.database.users_handler.get_user_by_id(
  769. user_id
  770. )
  771. if (
  772. user.limits_overrides
  773. and "max_collections" in user.limits_overrides
  774. ):
  775. return user.limits_overrides["max_collections"]
  776. return self.config.app.default_max_collections_per_user
  777. async def get_max_upload_size_by_type(
  778. self, user_id: UUID, file_type_or_ext: str
  779. ) -> int:
  780. """Return the maximum allowed upload size (in bytes) for the given
  781. user's file type/extension. Respects user-level overrides if present,
  782. falling back to the system config.
  783. ```json
  784. {
  785. "limits_overrides": {
  786. "max_file_size": 20_000_000,
  787. "max_file_size_by_type":
  788. {
  789. "pdf": 50_000_000,
  790. "docx": 30_000_000
  791. },
  792. ...
  793. }
  794. }
  795. ```
  796. """
  797. # 1. Normalize extension
  798. ext = file_type_or_ext.lower().lstrip(".")
  799. # 2. Fetch user from DB to see if we have any overrides
  800. user = await self.providers.database.users_handler.get_user_by_id(
  801. user_id
  802. )
  803. user_overrides = user.limits_overrides or {}
  804. # 3. Check if there's a user-level override for "max_file_size_by_type"
  805. user_file_type_limits = user_overrides.get("max_file_size_by_type", {})
  806. if ext in user_file_type_limits:
  807. return user_file_type_limits[ext]
  808. # 4. If not, check if there's a user-level fallback "max_file_size"
  809. if "max_file_size" in user_overrides:
  810. return user_overrides["max_file_size"]
  811. # 5. If none exist at user level, use system config
  812. # Example config paths:
  813. system_type_limits = self.config.app.max_upload_size_by_type
  814. if ext in system_type_limits:
  815. return system_type_limits[ext]
  816. # 6. Otherwise, return the global default
  817. return self.config.app.default_max_upload_size
  818. async def get_all_user_limits(self, user_id: UUID) -> dict[str, Any]:
  819. """
  820. Return a dictionary containing:
  821. - The system default limits (from self.config.limits)
  822. - The user's overrides (from user.limits_overrides)
  823. - The final 'effective' set of limits after merging (overall)
  824. - The usage for each relevant limit (per-route usage, etc.)
  825. """
  826. # 1) Fetch the user
  827. user = await self.providers.database.users_handler.get_user_by_id(
  828. user_id
  829. )
  830. user_overrides = user.limits_overrides or {}
  831. # 2) Grab system defaults
  832. system_defaults = {
  833. "global_per_min": self.config.database.limits.global_per_min,
  834. "route_per_min": self.config.database.limits.route_per_min,
  835. "monthly_limit": self.config.database.limits.monthly_limit,
  836. # Add additional fields if your LimitSettings has them
  837. }
  838. # 3) Build the overall (global) "effective limits" ignoring any specific route
  839. overall_effective = (
  840. self.providers.database.limits_handler.determine_effective_limits(
  841. user, route=""
  842. )
  843. )
  844. # 4) Build usage data. We'll do top-level usage for global_per_min/monthly,
  845. # then do route-by-route usage in a loop.
  846. usage: dict[str, Any] = {}
  847. now = datetime.now(timezone.utc)
  848. one_min_ago = now - timedelta(minutes=1)
  849. # (a) Global usage (per-minute)
  850. global_per_min_used = (
  851. await self.providers.database.limits_handler._count_requests(
  852. user_id, route=None, since=one_min_ago
  853. )
  854. )
  855. # (a2) Global usage (monthly) - i.e. usage across ALL routes
  856. global_monthly_used = await self.providers.database.limits_handler._count_monthly_requests(
  857. user_id, route=None
  858. )
  859. usage["global_per_min"] = {
  860. "used": global_per_min_used,
  861. "limit": overall_effective.global_per_min,
  862. "remaining": (
  863. overall_effective.global_per_min - global_per_min_used
  864. if overall_effective.global_per_min is not None
  865. else None
  866. ),
  867. }
  868. usage["monthly_limit"] = {
  869. "used": global_monthly_used,
  870. "limit": overall_effective.monthly_limit,
  871. "remaining": (
  872. overall_effective.monthly_limit - global_monthly_used
  873. if overall_effective.monthly_limit is not None
  874. else None
  875. ),
  876. }
  877. # (b) Route-level usage. We'll gather all routes from system + user overrides
  878. system_route_limits = (
  879. self.config.database.route_limits
  880. ) # dict[str, LimitSettings]
  881. user_route_overrides = user_overrides.get("route_overrides", {})
  882. route_keys = set(system_route_limits.keys()) | set(
  883. user_route_overrides.keys()
  884. )
  885. usage["routes"] = {}
  886. for route in route_keys:
  887. # 1) Get the final merged limits for this specific route
  888. route_effective = self.providers.database.limits_handler.determine_effective_limits(
  889. user, route
  890. )
  891. # 2) Count requests for the last minute on this route
  892. route_per_min_used = (
  893. await self.providers.database.limits_handler._count_requests(
  894. user_id, route, one_min_ago
  895. )
  896. )
  897. # 3) Count route-specific monthly usage
  898. route_monthly_used = await self.providers.database.limits_handler._count_monthly_requests(
  899. user_id, route
  900. )
  901. usage["routes"][route] = {
  902. "route_per_min": {
  903. "used": route_per_min_used,
  904. "limit": route_effective.route_per_min,
  905. "remaining": (
  906. route_effective.route_per_min - route_per_min_used
  907. if route_effective.route_per_min is not None
  908. else None
  909. ),
  910. },
  911. "monthly_limit": {
  912. "used": route_monthly_used,
  913. "limit": route_effective.monthly_limit,
  914. "remaining": (
  915. route_effective.monthly_limit - route_monthly_used
  916. if route_effective.monthly_limit is not None
  917. else None
  918. ),
  919. },
  920. }
  921. max_documents = await self.get_user_max_documents(user_id)
  922. used_documents = 0
  923. '''
  924. (
  925. await self.providers.database.documents_handler.get_documents_overview(
  926. limit=1, offset=0, filter_user_ids=[user_id]
  927. )
  928. )["total_entries"]
  929. '''
  930. max_chunks = await self.get_user_max_chunks(user_id)
  931. used_chunks = 0
  932. '''
  933. (
  934. await self.providers.database.chunks_handler.list_chunks(
  935. limit=1, offset=0, filters={"owner_id": user_id}
  936. )
  937. )["total_entries"]
  938. '''
  939. max_collections = await self.get_user_max_collections(user_id)
  940. used_collections: int = 0
  941. '''
  942. ( # type: ignore
  943. await self.providers.database.collections_handler.get_collections_overview(
  944. limit=1, offset=0, filter_user_ids=[user_id]
  945. )
  946. )["total_entries"]
  947. '''
  948. storage_limits = {
  949. "chunks": {
  950. "limit": max_chunks,
  951. "used": used_chunks,
  952. "remaining": (
  953. max_chunks - used_chunks
  954. if max_chunks is not None
  955. else None
  956. ),
  957. },
  958. "documents": {
  959. "limit": max_documents,
  960. "used": used_documents,
  961. "remaining": (
  962. max_documents - used_documents
  963. if max_documents is not None
  964. else None
  965. ),
  966. },
  967. "collections": {
  968. "limit": max_collections,
  969. "used": used_collections,
  970. "remaining": (
  971. max_collections - used_collections
  972. if max_collections is not None
  973. else None
  974. ),
  975. },
  976. }
  977. # 5) Return a structured response
  978. return {
  979. "storage_limits": storage_limits,
  980. "system_defaults": system_defaults,
  981. "user_overrides": user_overrides,
  982. "effective_limits": {
  983. "global_per_min": overall_effective.global_per_min,
  984. "route_per_min": overall_effective.route_per_min,
  985. "monthly_limit": overall_effective.monthly_limit,
  986. },
  987. "usage": usage,
  988. }