Coverage for src/lexigram/admin/data/adapters/repository/data_source.py: 0%

69 statements  

« prev     ^ index     » next       coverage.py v7.15.4, created at 2026-08-24 23:18 +0800

1"""AdminRepositoryProtocol-backed IDataSource adapter.""" 

2 

3from __future__ import annotations 

4 

5from typing import Any, Generic 

6 

7from lexigram.admin.data.adapters.repository.types import T 

8from lexigram.admin.data.data_source import IDataSource, QueryResult 

9from lexigram.admin.data.query import QuerySpec 

10from lexigram.contracts.admin.repository import AdminRepositoryProtocol 

11from lexigram.di.decorators import inject 

12from lexigram.logging import get_logger 

13 

14logger = get_logger(__name__) 

15 

16 

17@inject 

18class RepositoryDataSource(IDataSource[T], Generic[T]): 

19 """IDataSource adapter for AdminRepositoryProtocol repositories. 

20 

21 Translates QuerySpec into AdminRepositoryProtocol calls using 

22 QuerySpec.resolved_sort and QuerySpec.to_repository_filters(). 

23 """ 

24 

25 def __init__( 

26 self, 

27 repository: AdminRepositoryProtocol[T], 

28 *, 

29 tenant_scope: str | None = None, 

30 ) -> None: 

31 self._repo = repository 

32 # When set, every operation is constrained to this tenant: reads get a 

33 # mandatory ``tenant_id`` filter / post-fetch check, writes stamp or 

34 # refuse. See spec-security-remediation finding 3. 

35 self.tenant_scope = tenant_scope 

36 

37 def _scoped(self, filters: dict[str, Any] | None) -> dict[str, Any]: 

38 merged = dict(filters) if filters else {} 

39 if self.tenant_scope is not None: 

40 merged["tenant_id"] = self.tenant_scope 

41 return merged 

42 

43 def _in_scope(self, item: Any) -> bool: 

44 if self.tenant_scope is None: 

45 return True 

46 return str(getattr(item, "tenant_id", "")) == self.tenant_scope 

47 

48 async def find_one(self, item_id: Any) -> T | None: 

49 item = await self._repo.find_by_id(item_id) 

50 if item is not None and not self._in_scope(item): 

51 return None 

52 return item 

53 

54 async def find_many(self, query: QuerySpec) -> QueryResult[T]: 

55 sort_field, sort_order = query.resolved_sort 

56 filters = self._scoped(query.to_repository_filters()) 

57 

58 items = await self._repo.find_many( 

59 offset=query.offset, 

60 limit=query.per_page, 

61 order_by=[(sort_field, sort_order)] if sort_field else None, 

62 filters=filters or None, 

63 search=query.search or None, 

64 search_fields=query.search_fields or None, 

65 load=query.include or None, 

66 ) 

67 

68 total = await self.count(query) 

69 

70 return QueryResult( 

71 items=items, 

72 total=total, 

73 page=query.page, 

74 per_page=query.per_page, 

75 has_next=(query.offset + query.per_page) < total, 

76 has_prev=query.page > 1, 

77 ) 

78 

79 async def count(self, query: QuerySpec) -> int: 

80 return await self._repo.count( 

81 filters=self._scoped(query.to_repository_filters()), 

82 search=query.search or None, 

83 search_fields=query.search_fields or None, 

84 ) 

85 

86 async def create(self, data: dict[str, Any]) -> T: 

87 if self.tenant_scope is not None: 

88 data = {**data, "tenant_id": self.tenant_scope} 

89 return await self._repo.create(data) 

90 

91 async def update(self, item_id: Any, data: dict[str, Any]) -> T: 

92 if self.tenant_scope is not None: 

93 existing = await self._repo.find_by_id(item_id) 

94 if existing is not None and not self._in_scope(existing): 

95 raise PermissionError(f"item {item_id!r} belongs to another tenant") 

96 return await self._repo.update(item_id, data) 

97 

98 async def delete(self, item_id: Any) -> bool: 

99 if self.tenant_scope is not None: 

100 existing = await self._repo.find_by_id(item_id) 

101 if existing is not None and not self._in_scope(existing): 

102 raise PermissionError(f"item {item_id!r} belongs to another tenant") 

103 return await self._repo.delete(item_id) 

104 

105 async def bulk_create(self, items: list[dict[str, Any]]) -> list[T]: 

106 return [await self._repo.create(item) for item in items] 

107 

108 async def bulk_update(self, ids: list[Any], data: dict[str, Any]) -> int: 

109 count = 0 

110 for item_id in ids: 

111 try: 

112 await self._repo.update(item_id, data) 

113 count += 1 

114 except Exception: 

115 logger.warning("bulk_update: skipped id=%s", item_id) 

116 return count 

117 

118 async def bulk_delete(self, ids: list[Any]) -> int: 

119 count = 0 

120 for item_id in ids: 

121 if await self._repo.delete(item_id): 

122 count += 1 

123 return count