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
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 02:04 +0000
1from __future__ import annotations
3import base64
4from typing import TYPE_CHECKING, Any, Dict, List, Optional
5from urllib.parse import quote
6from uuid import UUID
8import httpx
9import pydantic
10from httpx import Response
11from starlette import status
12from typing_extensions import Self
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
24if TYPE_CHECKING: 24 ↛ 25line 24 didn't jump to line 25 because the condition on line 24 was never true
25 import logging
27logger: "logging.Logger" = get_logger(__name__)
30class BaseClient:
31 _http_client: PrefectHttpxAsyncClient
33 def __init__(self, additional_headers: dict[str, str] | None = None):
34 from prefect.server.api.server import create_app
36 additional_headers = additional_headers or {}
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()
42 settings = get_current_settings()
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}")
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 )
58 async def __aenter__(self) -> Self:
59 await self._http_client.__aenter__()
60 return self
62 async def __aexit__(self, *args: Any) -> None:
63 await self._http_client.__aexit__(*args)
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}")
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())
82 async def read_flow_raw(self, flow_id: UUID) -> Response:
83 return await self._http_client.get(f"/flows/{flow_id}")
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 )
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}")
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}")
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}")
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())
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 )
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 )
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 )
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 )
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 )
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 )
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()
150 pools = pydantic.TypeAdapter(List[WorkPool]).validate_python(response.json())
151 return pools[0] if pools else None
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}")
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")
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 )
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 )
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 )
181 VARIABLE_PAGE_SIZE = 200
182 MAX_VARIABLES_PER_WORKSPACE = 1000
184 async def read_workspace_variables(
185 self, names: Optional[List[str]] = None
186 ) -> Dict[str, StrictVariableValue]:
187 variables: Dict[str, StrictVariableValue] = {}
189 offset = 0
191 filter = VariableFilter()
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)))
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()
212 results = response.json()
213 for variable in results:
214 variables[variable["name"]] = variable["value"]
216 if len(results) < self.VARIABLE_PAGE_SIZE:
217 break
219 return variables
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 )
229class WorkPoolsOrchestrationClient(BaseClient):
230 async def __aenter__(self) -> Self:
231 return self
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