Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 2 additions & 2 deletions .github/workflows/ci.yml
Original file line number Diff line number Diff line change
Expand Up @@ -24,9 +24,9 @@ jobs:
- name: Install dependencies
run: pip install -e ".[dev]"
- name: Ruff check
run: ruff check src/ tests/ scripts/
run: ruff check src/ tests/ scripts/ sdk/
- name: Ruff format check
run: ruff format --check src/ tests/ scripts/
run: ruff format --check src/ tests/ scripts/ sdk/
- name: Type check
run: mypy src/ --ignore-missing-imports

Expand Down
8 changes: 4 additions & 4 deletions Makefile
Original file line number Diff line number Diff line change
Expand Up @@ -105,13 +105,13 @@ perf-plot:
# ── Code Quality ──────────────────────────────────────────────────

lint:
ruff check src/ tests/
ruff format --check src/ tests/
ruff check src/ tests/ scripts/ sdk/
ruff format --check src/ tests/ scripts/ sdk/
mypy src/

format:
ruff format src/ tests/
ruff check --fix src/ tests/
ruff format src/ tests/ scripts/ sdk/
ruff check --fix src/ tests/ scripts/ sdk/

# ── Build & Deploy ────────────────────────────────────────────────

Expand Down
50 changes: 10 additions & 40 deletions sdk/agentflow/async_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -197,9 +197,7 @@ def _record_version_headers(self, headers: httpx.Headers) -> None:
self._last_server_version = headers.get("X-AgentFlow-Version")
self._last_latest_version = headers.get("X-AgentFlow-Latest-Version")
self._last_deprecated = headers.get("X-AgentFlow-Deprecated")
self._last_deprecation_warning = headers.get(
"X-AgentFlow-Deprecation-Warning"
)
self._last_deprecation_warning = headers.get("X-AgentFlow-Deprecation-Warning")

async def _get_entity(
self,
Expand All @@ -219,9 +217,7 @@ def _parse_contract_versions(
return {}
entity, separator, version = contract_version.partition(":")
if not separator or not entity or not version:
raise ValueError(
"contract_version must use '<entity>:<version>' format."
)
raise ValueError("contract_version must use '<entity>:<version>' format.")
return {entity: version[1:] if version.startswith("v") else version}

async def _apply_contract_version(
Expand All @@ -234,27 +230,14 @@ async def _apply_contract_version(
return payload
contract = await self._get_contract(entity_type, version)
fields = contract.get("fields", [])
required_fields = [
field["name"]
for field in fields
if field.get("required")
]
missing_fields = [
field_name
for field_name in required_fields
if field_name not in payload
]
required_fields = [field["name"] for field in fields if field.get("required")]
missing_fields = [field_name for field_name in required_fields if field_name not in payload]
if missing_fields:
raise AgentFlowError(
"Contract validation failed. Missing required fields: "
+ ", ".join(missing_fields)
"Contract validation failed. Missing required fields: " + ", ".join(missing_fields)
)
allowed_fields = {field["name"] for field in fields}
return {
name: value
for name, value in payload.items()
if name in allowed_fields
}
return {name: value for name, value in payload.items() if name in allowed_fields}

async def _get_contract(self, entity_type: str, version: str) -> dict[str, Any]:
cache_key = (entity_type, version)
Expand Down Expand Up @@ -355,11 +338,7 @@ async def _query_page(
payload["limit"] = limit
if cursor is not None:
payload["cursor"] = cursor
headers = (
{"Idempotency-Key": idempotency_key}
if idempotency_key is not None
else None
)
headers = {"Idempotency-Key": idempotency_key} if idempotency_key is not None else None
return await self._request("POST", "/v1/query", json=payload, headers=headers)

async def query(
Expand Down Expand Up @@ -404,8 +383,7 @@ async def search(
async def list_contracts(self) -> list[ContractSummary]:
payload = await self._request("GET", "/v1/contracts")
return [
ContractSummary.model_validate(contract)
for contract in payload.get("contracts", [])
ContractSummary.model_validate(contract) for contract in payload.get("contracts", [])
]

async def get_contract(
Expand Down Expand Up @@ -438,11 +416,7 @@ async def validate_contract(
*,
idempotency_key: str | None = None,
) -> ContractValidation:
headers = (
{"Idempotency-Key": idempotency_key}
if idempotency_key is not None
else None
)
headers = {"Idempotency-Key": idempotency_key} if idempotency_key is not None else None
response = await self._request(
"POST",
f"/v1/contracts/{entity}/validate",
Expand Down Expand Up @@ -500,11 +474,7 @@ async def batch(
*,
idempotency_key: str | None = None,
) -> dict[str, Any]:
headers = (
{"Idempotency-Key": idempotency_key}
if idempotency_key is not None
else None
)
headers = {"Idempotency-Key": idempotency_key} if idempotency_key is not None else None
return await self._request(
"POST",
"/v1/batch",
Expand Down
5 changes: 1 addition & 4 deletions sdk/agentflow/cli.py
Original file line number Diff line number Diff line change
Expand Up @@ -124,10 +124,7 @@ def _scaffold_project(
continue
relative_path = source_path.relative_to(template_dir)
target_path = project_dir.joinpath(
*[
part[:-5] if part.endswith(".tmpl") else part
for part in relative_path.parts
]
*[part[:-5] if part.endswith(".tmpl") else part for part in relative_path.parts]
)
target_path.parent.mkdir(parents=True, exist_ok=True)
rendered = _render_template(
Expand Down
50 changes: 10 additions & 40 deletions sdk/agentflow/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -197,9 +197,7 @@ def _record_version_headers(self, headers: httpx.Headers) -> None:
self._last_server_version = headers.get("X-AgentFlow-Version")
self._last_latest_version = headers.get("X-AgentFlow-Latest-Version")
self._last_deprecated = headers.get("X-AgentFlow-Deprecated")
self._last_deprecation_warning = headers.get(
"X-AgentFlow-Deprecation-Warning"
)
self._last_deprecation_warning = headers.get("X-AgentFlow-Deprecation-Warning")

def _get_entity(
self,
Expand All @@ -222,9 +220,7 @@ def _parse_contract_versions(
return {}
entity, separator, version = contract_version.partition(":")
if not separator or not entity or not version:
raise ValueError(
"contract_version must use '<entity>:<version>' format."
)
raise ValueError("contract_version must use '<entity>:<version>' format.")
return {entity: version[1:] if version.startswith("v") else version}

def _apply_contract_version(
Expand All @@ -237,27 +233,14 @@ def _apply_contract_version(
return payload
contract = self._get_contract(entity_type, version)
fields = contract.get("fields", [])
required_fields = [
field["name"]
for field in fields
if field.get("required")
]
missing_fields = [
field_name
for field_name in required_fields
if field_name not in payload
]
required_fields = [field["name"] for field in fields if field.get("required")]
missing_fields = [field_name for field_name in required_fields if field_name not in payload]
if missing_fields:
raise AgentFlowError(
"Contract validation failed. Missing required fields: "
+ ", ".join(missing_fields)
"Contract validation failed. Missing required fields: " + ", ".join(missing_fields)
)
allowed_fields = {field["name"] for field in fields}
return {
name: value
for name, value in payload.items()
if name in allowed_fields
}
return {name: value for name, value in payload.items() if name in allowed_fields}

def _get_contract(self, entity_type: str, version: str) -> dict[str, Any]:
cache_key = (entity_type, version)
Expand Down Expand Up @@ -359,11 +342,7 @@ def _query_page(
payload["limit"] = limit
if cursor is not None:
payload["cursor"] = cursor
headers = (
{"Idempotency-Key": idempotency_key}
if idempotency_key is not None
else None
)
headers = {"Idempotency-Key": idempotency_key} if idempotency_key is not None else None
return self._request("POST", "/v1/query", json=payload, headers=headers)

def query(
Expand Down Expand Up @@ -408,8 +387,7 @@ def search(
def list_contracts(self) -> list[ContractSummary]:
payload = self._request("GET", "/v1/contracts")
return [
ContractSummary.model_validate(contract)
for contract in payload.get("contracts", [])
ContractSummary.model_validate(contract) for contract in payload.get("contracts", [])
]

def get_contract(
Expand Down Expand Up @@ -442,11 +420,7 @@ def validate_contract(
*,
idempotency_key: str | None = None,
) -> ContractValidation:
headers = (
{"Idempotency-Key": idempotency_key}
if idempotency_key is not None
else None
)
headers = {"Idempotency-Key": idempotency_key} if idempotency_key is not None else None
response = self._request(
"POST",
f"/v1/contracts/{entity}/validate",
Expand Down Expand Up @@ -504,11 +478,7 @@ def batch(
*,
idempotency_key: str | None = None,
) -> dict[str, Any]:
headers = (
{"Idempotency-Key": idempotency_key}
if idempotency_key is not None
else None
)
headers = {"Idempotency-Key": idempotency_key} if idempotency_key is not None else None
return self._request(
"POST",
"/v1/batch",
Expand Down