Coverage for /usr/local/lib/python3.12/site-packages/prefect/server/api/clients.py: 45%

126 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-10-07 02:04 +0000

1from __future__ import annotations 

2 

3import base64 

4from typing import TYPE_CHECKING, Any, Dict, List, Optional 

5from urllib.parse import quote 

6from uuid import UUID 

7 

8import httpx 

9import pydantic 

10from httpx import Response 

11from starlette import status 

12from typing_extensions import Self 

13 

14from prefect.client.base import PrefectHttpxAsyncClient 

15from prefect.exceptions import ObjectNotFound 

16from prefect.logging import get_logger 

17from prefect.server.schemas.actions import DeploymentFlowRunCreate, StateCreate 

18from prefect.server.schemas.core import WorkPool 

19from prefect.server.schemas.filters import VariableFilter, VariableFilterName 

20from prefect.server.schemas.responses import DeploymentResponse, OrchestrationResult 

21from prefect.settings import get_current_settings 

22from prefect.types import StrictVariableValue 

23 

24if TYPE_CHECKING: 24 ↛ 25line 24 didn't jump to line 25 because the condition on line 24 was never true

25 import logging 

26 

27logger: "logging.Logger" = get_logger(__name__) 

28 

29 

30class BaseClient: 

31 _http_client: PrefectHttpxAsyncClient 

32 

33 def __init__(self, additional_headers: dict[str, str] | None = None): 

34 from prefect.server.api.server import create_app 

35 

36 additional_headers = additional_headers or {} 

37 

38 # create_app caches application instances, and invoking it with no arguments 

39 # will point it to the the currently running server instance 

40 api_app = create_app() 

41 

42 settings = get_current_settings() 

43 

44 # we pull the auth string from _server_ settings because this client is run on the server 

45 if auth_string_secret := settings.server.api.auth_string: 45 ↛ 46line 45 didn't jump to line 46 because the condition on line 45 was never true

46 if auth_string := auth_string_secret.get_secret_value(): 

47 token = base64.b64encode(auth_string.encode("utf-8")).decode("utf-8") 

48 additional_headers.setdefault("Authorization", f"Basic {token}") 

49 

50 self._http_client = PrefectHttpxAsyncClient( 

51 transport=httpx.ASGITransport(app=api_app, raise_app_exceptions=False), 

52 headers={**additional_headers}, 

53 base_url=f"http://prefect-in-memory{settings.server.api.base_path or '/api'}", 

54 enable_csrf_support=settings.server.api.csrf_protection_enabled, 

55 raise_on_all_errors=False, 

56 ) 

57 

58 async def __aenter__(self) -> Self: 

59 await self._http_client.__aenter__() 

60 return self 

61 

62 async def __aexit__(self, *args: Any) -> None: 

63 await self._http_client.__aexit__(*args) 

64 

65 

66class OrchestrationClient(BaseClient): 

67 async def read_deployment_raw(self, deployment_id: UUID) -> Response: 

68 return await self._http_client.get(f"/deployments/{deployment_id}") 

69 

70 async def read_deployment( 

71 self, deployment_id: UUID 

72 ) -> Optional[DeploymentResponse]: 

73 try: 

74 response = await self.read_deployment_raw(deployment_id) 

75 response.raise_for_status() 

76 except httpx.HTTPStatusError as e: 

77 if e.response.status_code == status.HTTP_404_NOT_FOUND: 

78 return None 

79 raise 

80 return DeploymentResponse.model_validate(response.json()) 

81 

82 async def read_flow_raw(self, flow_id: UUID) -> Response: 

83 return await self._http_client.get(f"/flows/{flow_id}") 

84 

85 async def create_flow_run( 

86 self, deployment_id: UUID, flow_run_create: DeploymentFlowRunCreate 

87 ) -> Response: 

88 return await self._http_client.post( 

89 f"/deployments/{deployment_id}/create_flow_run", 

90 json=flow_run_create.model_dump(mode="json"), 

91 ) 

92 

93 async def read_flow_run_raw(self, flow_run_id: UUID) -> Response: 

94 return await self._http_client.get(f"/flow_runs/{flow_run_id}") 

95 

96 async def delete_flow_run(self, flow_run_id: UUID) -> Response: 

97 return await self._http_client.delete(f"/flow_runs/{flow_run_id}") 

98 

99 async def read_task_run_raw(self, task_run_id: UUID) -> Response: 

100 return await self._http_client.get(f"/task_runs/{task_run_id}") 

101 

102 async def resume_flow_run(self, flow_run_id: UUID) -> OrchestrationResult: 

103 response = await self._http_client.post( 

104 f"/flow_runs/{flow_run_id}/resume", 

105 ) 

106 response.raise_for_status() 

107 return OrchestrationResult.model_validate(response.json()) 

108 

109 async def pause_deployment(self, deployment_id: UUID) -> Response: 

110 return await self._http_client.post( 

111 f"/deployments/{deployment_id}/pause_deployment", 

112 ) 

113 

114 async def resume_deployment(self, deployment_id: UUID) -> Response: 

115 return await self._http_client.post( 

116 f"/deployments/{deployment_id}/resume_deployment", 

117 ) 

118 

119 async def set_flow_run_state( 

120 self, flow_run_id: UUID, state: StateCreate, force: bool = False 

121 ) -> Response: 

122 return await self._http_client.post( 

123 f"/flow_runs/{flow_run_id}/set_state", 

124 json={ 

125 "state": state.model_dump(mode="json"), 

126 "force": force, 

127 }, 

128 ) 

129 

130 async def pause_work_pool(self, work_pool_name: str) -> Response: 

131 return await self._http_client.patch( 

132 f"/work_pools/{quote(work_pool_name)}", json={"is_paused": True} 

133 ) 

134 

135 async def resume_work_pool(self, work_pool_name: str) -> Response: 

136 return await self._http_client.patch( 

137 f"/work_pools/{quote(work_pool_name)}", json={"is_paused": False} 

138 ) 

139 

140 async def read_work_pool_raw(self, work_pool_id: UUID) -> Response: 

141 return await self._http_client.post( 

142 "/work_pools/filter", 

143 json={"work_pools": {"id": {"any_": [str(work_pool_id)]}}}, 

144 ) 

145 

146 async def read_work_pool(self, work_pool_id: UUID) -> Optional[WorkPool]: 

147 response = await self.read_work_pool_raw(work_pool_id) 

148 response.raise_for_status() 

149 

150 pools = pydantic.TypeAdapter(List[WorkPool]).validate_python(response.json()) 

151 return pools[0] if pools else None 

152 

153 async def read_work_queue_raw(self, work_queue_id: UUID) -> Response: 

154 return await self._http_client.get(f"/work_queues/{work_queue_id}") 

155 

156 async def read_work_queue_status_raw(self, work_queue_id: UUID) -> Response: 

157 return await self._http_client.get(f"/work_queues/{work_queue_id}/status") 

158 

159 async def pause_work_queue(self, work_queue_id: UUID) -> Response: 

160 return await self._http_client.patch( 

161 f"/work_queues/{work_queue_id}", 

162 json={"is_paused": True}, 

163 ) 

164 

165 async def resume_work_queue(self, work_queue_id: UUID) -> Response: 

166 return await self._http_client.patch( 

167 f"/work_queues/{work_queue_id}", 

168 json={"is_paused": False}, 

169 ) 

170 

171 async def read_block_document_raw( 

172 self, 

173 block_document_id: UUID, 

174 include_secrets: bool = True, 

175 ) -> Response: 

176 return await self._http_client.get( 

177 f"/block_documents/{block_document_id}", 

178 params=dict(include_secrets=include_secrets), 

179 ) 

180 

181 VARIABLE_PAGE_SIZE = 200 

182 MAX_VARIABLES_PER_WORKSPACE = 1000 

183 

184 async def read_workspace_variables( 

185 self, names: Optional[List[str]] = None 

186 ) -> Dict[str, StrictVariableValue]: 

187 variables: Dict[str, StrictVariableValue] = {} 

188 

189 offset = 0 

190 

191 filter = VariableFilter() 

192 

193 if names is not None and not names: 

194 return variables 

195 elif names is not None: 

196 filter.name = VariableFilterName(any_=list(set(names))) 

197 

198 for offset in range( 

199 0, self.MAX_VARIABLES_PER_WORKSPACE, self.VARIABLE_PAGE_SIZE 

200 ): 

201 response = await self._http_client.post( 

202 "/variables/filter", 

203 json={ 

204 "variables": filter.model_dump(), 

205 "limit": self.VARIABLE_PAGE_SIZE, 

206 "offset": offset, 

207 }, 

208 ) 

209 if response.status_code >= 300: 

210 response.raise_for_status() 

211 

212 results = response.json() 

213 for variable in results: 

214 variables[variable["name"]] = variable["value"] 

215 

216 if len(results) < self.VARIABLE_PAGE_SIZE: 

217 break 

218 

219 return variables 

220 

221 async def read_concurrency_limit_v2_raw( 

222 self, concurrency_limit_id: UUID 

223 ) -> Response: 

224 return await self._http_client.get( 

225 f"/v2/concurrency_limits/{concurrency_limit_id}" 

226 ) 

227 

228 

229class WorkPoolsOrchestrationClient(BaseClient): 

230 async def __aenter__(self) -> Self: 

231 return self 

232 

233 async def read_work_pool(self, work_pool_name: str) -> WorkPool: 

234 """ 

235 Reads information for a given work pool 

236 Args: 

237 work_pool_name: The name of the work pool to for which to get 

238 information. 

239 Returns: 

240 Information about the requested work pool. 

241 """ 

242 try: 

243 response = await self._http_client.get(f"/work_pools/{work_pool_name}") 

244 response.raise_for_status() 

245 return WorkPool.model_validate(response.json()) 

246 except httpx.HTTPStatusError as e: 

247 if e.response.status_code == status.HTTP_404_NOT_FOUND: 

248 raise ObjectNotFound(http_exc=e) from e 

249 else: 

250 raise