mirror of
https://github.com/microsoft/agent-framework.git
synced 2026-06-16 21:04:09 +08:00
Python: Fix Http Schema (#2112)
* Rename to threadid * Respond in plain text * Make snake-case * Add http prefix * rename to wait-for-response * Add query param check * address comments
This commit is contained in:
committed by
GitHub
Unverified
parent
ebab25b196
commit
ff28066c9c
@@ -8,7 +8,7 @@ with Azure Durable Entities, enabling stateful and durable AI agent execution.
|
||||
|
||||
import json
|
||||
import re
|
||||
from collections.abc import Mapping
|
||||
from collections.abc import Callable, Mapping
|
||||
from typing import Any, cast
|
||||
|
||||
import azure.durable_functions as df
|
||||
@@ -18,17 +18,16 @@ from agent_framework import AgentProtocol, get_logger
|
||||
from ._callbacks import AgentResponseCallbackProtocol
|
||||
from ._entities import create_agent_entity
|
||||
from ._errors import IncomingRequestError
|
||||
from ._models import AgentSessionId, ChatRole, RunRequest
|
||||
from ._models import AgentSessionId, RunRequest
|
||||
from ._state import AgentState
|
||||
|
||||
logger = get_logger("agent_framework.azurefunctions")
|
||||
|
||||
SESSION_ID_FIELD: str = "sessionId"
|
||||
SESSION_KEY_FIELD: str = "sessionKey"
|
||||
SESSION_IDENTIFIER_KEYS: tuple[str, str] = (
|
||||
SESSION_ID_FIELD,
|
||||
SESSION_KEY_FIELD,
|
||||
)
|
||||
THREAD_ID_FIELD: str = "thread_id"
|
||||
RESPONSE_FORMAT_JSON: str = "json"
|
||||
RESPONSE_FORMAT_TEXT: str = "text"
|
||||
WAIT_FOR_RESPONSE_FIELD: str = "wait_for_response"
|
||||
WAIT_FOR_RESPONSE_HEADER: str = "x-ms-wait-for-response"
|
||||
|
||||
|
||||
class AgentFunctionApp(df.DFApp):
|
||||
@@ -88,18 +87,18 @@ class AgentFunctionApp(df.DFApp):
|
||||
def __init__(
|
||||
self,
|
||||
agents: list[AgentProtocol] | None = None,
|
||||
http_auth_level: func.AuthLevel = func.AuthLevel.ANONYMOUS,
|
||||
http_auth_level: func.AuthLevel = func.AuthLevel.FUNCTION,
|
||||
enable_health_check: bool = True,
|
||||
enable_http_endpoints: bool = True,
|
||||
max_poll_retries: int = 10,
|
||||
poll_interval_seconds: float = 0.5,
|
||||
max_poll_retries: int = 30,
|
||||
poll_interval_seconds: float = 1,
|
||||
default_callback: AgentResponseCallbackProtocol | None = None,
|
||||
):
|
||||
"""Initialize the AgentFunctionApp.
|
||||
|
||||
Args:
|
||||
agents: List of agent instances to register
|
||||
http_auth_level: HTTP authentication level (default: ANONYMOUS)
|
||||
http_auth_level: HTTP authentication level (default: FUNCTION)
|
||||
enable_health_check: Enable built-in health check endpoint (default: True)
|
||||
enable_http_endpoints: Enable HTTP endpoints for agents (default: True)
|
||||
max_poll_retries: Maximum number of polling attempts when waiting for a response
|
||||
@@ -231,7 +230,7 @@ class AgentFunctionApp(df.DFApp):
|
||||
Args:
|
||||
agent_name: The agent name (used for both routing and entity identification)
|
||||
"""
|
||||
run_function_name = self._build_function_name(agent_name, "run")
|
||||
run_function_name = self._build_function_name(agent_name, "http")
|
||||
|
||||
@self.function_name(run_function_name)
|
||||
@self.route(route=f"agents/{agent_name}/run", methods=["POST"])
|
||||
@@ -242,7 +241,7 @@ class AgentFunctionApp(df.DFApp):
|
||||
Expected request body (RunRequest format):
|
||||
{
|
||||
"message": "user message to agent",
|
||||
"sessionId": "optional session id (or sessionKey)",
|
||||
"thread_id": "optional conversation identifier",
|
||||
"role": "user|system" (optional, default: "user"),
|
||||
"response_format": {...} (optional JSON schema for structured responses),
|
||||
"enable_tool_calls": true|false (optional, default: true)
|
||||
@@ -250,22 +249,28 @@ class AgentFunctionApp(df.DFApp):
|
||||
"""
|
||||
logger.debug(f"[HTTP Trigger] Received request on route: /api/agents/{agent_name}/run")
|
||||
|
||||
response_format: str = RESPONSE_FORMAT_JSON
|
||||
thread_id: str | None = None
|
||||
|
||||
try:
|
||||
req_body, message = self._parse_incoming_request(req)
|
||||
session_key = self._resolve_session_key(req=req, req_body=req_body)
|
||||
wait_for_completion = self._should_wait_for_completion(req=req, req_body=req_body)
|
||||
req_body, message, response_format = self._parse_incoming_request(req)
|
||||
thread_id = self._resolve_thread_id(req=req, req_body=req_body)
|
||||
wait_for_response = self._should_wait_for_response(req=req, req_body=req_body)
|
||||
|
||||
logger.debug(f"[HTTP Trigger] Message: {message}")
|
||||
logger.debug(f"[HTTP Trigger] Session Key: {session_key}")
|
||||
logger.debug(f"[HTTP Trigger] wait_for_completion: {wait_for_completion}")
|
||||
logger.debug(f"[HTTP Trigger] Thread ID: {thread_id}")
|
||||
logger.debug(f"[HTTP Trigger] wait_for_response: {wait_for_response}")
|
||||
|
||||
if not message:
|
||||
logger.warning("[HTTP Trigger] Request rejected: Missing message")
|
||||
return func.HttpResponse(
|
||||
json.dumps({"error": "Message is required"}), status_code=400, mimetype="application/json"
|
||||
return self._create_http_response(
|
||||
payload={"error": "Message is required"},
|
||||
status_code=400,
|
||||
response_format=response_format,
|
||||
thread_id=thread_id,
|
||||
)
|
||||
|
||||
session_id = self._create_session_id(agent_name, session_key)
|
||||
session_id = self._create_session_id(agent_name, thread_id)
|
||||
correlation_id = self._generate_unique_id()
|
||||
|
||||
logger.debug(f"[HTTP Trigger] Using session ID: {session_id}")
|
||||
@@ -276,7 +281,7 @@ class AgentFunctionApp(df.DFApp):
|
||||
run_request = self._build_request_data(
|
||||
req_body,
|
||||
message,
|
||||
session_key,
|
||||
thread_id,
|
||||
correlation_id,
|
||||
)
|
||||
logger.debug("Signalling entity %s with request: %s", entity_instance_id, run_request)
|
||||
@@ -284,43 +289,60 @@ class AgentFunctionApp(df.DFApp):
|
||||
|
||||
logger.debug(f"[HTTP Trigger] Signal sent to entity {session_id}")
|
||||
|
||||
if wait_for_completion:
|
||||
if wait_for_response:
|
||||
result = await self._get_response_from_entity(
|
||||
client=client,
|
||||
entity_instance_id=entity_instance_id,
|
||||
correlation_id=correlation_id,
|
||||
message=message,
|
||||
session_key=session_key,
|
||||
thread_id=thread_id,
|
||||
)
|
||||
|
||||
logger.debug(f"[HTTP Trigger] Result status: {result.get('status', 'unknown')}")
|
||||
return func.HttpResponse(
|
||||
json.dumps(result),
|
||||
return self._create_http_response(
|
||||
payload=result,
|
||||
status_code=200 if result.get("status") == "success" else 500,
|
||||
mimetype="application/json",
|
||||
response_format=response_format,
|
||||
thread_id=thread_id,
|
||||
)
|
||||
|
||||
logger.debug("[HTTP Trigger] wait_for_completion disabled; returning correlation ID")
|
||||
logger.debug("[HTTP Trigger] wait_for_response disabled; returning correlation ID")
|
||||
|
||||
accepted_response = self._build_accepted_response(
|
||||
message=message, session_key=session_key, correlation_id=correlation_id
|
||||
message=message, thread_id=thread_id, correlation_id=correlation_id
|
||||
)
|
||||
|
||||
return func.HttpResponse(json.dumps(accepted_response), status_code=202, mimetype="application/json")
|
||||
return self._create_http_response(
|
||||
payload=accepted_response,
|
||||
status_code=202,
|
||||
response_format=response_format,
|
||||
thread_id=thread_id,
|
||||
)
|
||||
|
||||
except IncomingRequestError as exc:
|
||||
logger.warning(f"[HTTP Trigger] Request rejected: {exc!s}")
|
||||
return func.HttpResponse(
|
||||
json.dumps({"error": str(exc)}), status_code=exc.status_code, mimetype="application/json"
|
||||
return self._create_http_response(
|
||||
payload={"error": str(exc)},
|
||||
status_code=exc.status_code,
|
||||
response_format=response_format,
|
||||
thread_id=thread_id,
|
||||
)
|
||||
except ValueError as exc:
|
||||
logger.error(f"[HTTP Trigger] Invalid JSON: {exc!s}")
|
||||
return func.HttpResponse(
|
||||
json.dumps({"error": "Invalid JSON"}), status_code=400, mimetype="application/json"
|
||||
return self._create_http_response(
|
||||
payload={"error": "Invalid JSON"},
|
||||
status_code=400,
|
||||
response_format=response_format,
|
||||
thread_id=thread_id,
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.error(f"[HTTP Trigger] Error: {exc!s}", exc_info=True)
|
||||
return func.HttpResponse(json.dumps({"error": str(exc)}), status_code=500, mimetype="application/json")
|
||||
return self._create_http_response(
|
||||
payload={"error": str(exc)},
|
||||
status_code=500,
|
||||
response_format=response_format,
|
||||
thread_id=thread_id,
|
||||
)
|
||||
|
||||
_ = http_start
|
||||
|
||||
@@ -365,7 +387,7 @@ class AgentFunctionApp(df.DFApp):
|
||||
{
|
||||
"name": name,
|
||||
"type": type(agent).__name__,
|
||||
"httpEndpointEnabled": self.agent_http_endpoint_flags.get(
|
||||
"http_endpoint_enabled": self.agent_http_endpoint_flags.get(
|
||||
name,
|
||||
self.enable_http_endpoints,
|
||||
),
|
||||
@@ -381,17 +403,20 @@ class AgentFunctionApp(df.DFApp):
|
||||
_ = health_check
|
||||
|
||||
@staticmethod
|
||||
def _build_function_name(agent_name: str, suffix: str) -> str:
|
||||
"""Generate a unique, Azure Functions-compliant name for an agent function."""
|
||||
sanitized = re.sub(r"[^0-9a-zA-Z_]", "_", agent_name or "agent").strip("_")
|
||||
def _build_function_name(agent_name: str, prefix: str) -> str:
|
||||
"""Generate the sanitized function name in the form "{prefix}-{sanitized_agent_name}".
|
||||
|
||||
if not sanitized:
|
||||
sanitized = "agent"
|
||||
Example: agent_name="Weather Agent" and prefix="http" becomes "http-Weather_Agent".
|
||||
"""
|
||||
sanitized_agent = re.sub(r"[^0-9a-zA-Z_]", "_", agent_name or "agent").strip("_")
|
||||
|
||||
if sanitized[0].isdigit():
|
||||
sanitized = f"agent_{sanitized}"
|
||||
if not sanitized_agent:
|
||||
sanitized_agent = "agent"
|
||||
|
||||
return f"{sanitized}_{suffix}"
|
||||
if sanitized_agent[0].isdigit():
|
||||
sanitized_agent = f"agent_{sanitized_agent}"
|
||||
|
||||
return f"{prefix}-{sanitized_agent}"
|
||||
|
||||
async def _read_cached_state(
|
||||
self,
|
||||
@@ -418,7 +443,7 @@ class AgentFunctionApp(df.DFApp):
|
||||
entity_instance_id: df.EntityId,
|
||||
correlation_id: str,
|
||||
message: str,
|
||||
session_key: str,
|
||||
thread_id: str,
|
||||
) -> dict[str, Any]:
|
||||
"""Poll the entity state until a response is available or timeout occurs."""
|
||||
import asyncio
|
||||
@@ -438,7 +463,7 @@ class AgentFunctionApp(df.DFApp):
|
||||
entity_instance_id=entity_instance_id,
|
||||
correlation_id=correlation_id,
|
||||
message=message,
|
||||
session_key=session_key,
|
||||
thread_id=thread_id,
|
||||
)
|
||||
if result is not None:
|
||||
break
|
||||
@@ -453,7 +478,7 @@ class AgentFunctionApp(df.DFApp):
|
||||
f"[HTTP Trigger] Response with correlation ID {correlation_id} "
|
||||
f"not found in time (waited {max_retries * interval} seconds)"
|
||||
)
|
||||
return await self._build_timeout_result(message=message, session_key=session_key, correlation_id=correlation_id)
|
||||
return await self._build_timeout_result(message=message, thread_id=thread_id, correlation_id=correlation_id)
|
||||
|
||||
async def _poll_entity_for_response(
|
||||
self,
|
||||
@@ -461,7 +486,7 @@ class AgentFunctionApp(df.DFApp):
|
||||
entity_instance_id: df.EntityId,
|
||||
correlation_id: str,
|
||||
message: str,
|
||||
session_key: str,
|
||||
thread_id: str,
|
||||
) -> dict[str, Any] | None:
|
||||
result: dict[str, Any] | None = None
|
||||
try:
|
||||
@@ -475,7 +500,7 @@ class AgentFunctionApp(df.DFApp):
|
||||
result = self._build_success_result(
|
||||
response_data=agent_response,
|
||||
message=message,
|
||||
session_key=session_key,
|
||||
thread_id=thread_id,
|
||||
correlation_id=correlation_id,
|
||||
state=state,
|
||||
)
|
||||
@@ -486,113 +511,181 @@ class AgentFunctionApp(df.DFApp):
|
||||
|
||||
return result
|
||||
|
||||
async def _build_timeout_result(self, message: str, session_key: str, correlation_id: str) -> dict[str, Any]:
|
||||
async def _build_timeout_result(self, message: str, thread_id: str, correlation_id: str) -> dict[str, Any]:
|
||||
"""Create the timeout response."""
|
||||
return {
|
||||
"response": "Agent is still processing or timed out...",
|
||||
"message": message,
|
||||
SESSION_ID_FIELD: session_key,
|
||||
THREAD_ID_FIELD: thread_id,
|
||||
"status": "timeout",
|
||||
"correlationId": correlation_id,
|
||||
"correlation_id": correlation_id,
|
||||
}
|
||||
|
||||
def _build_success_result(
|
||||
self, response_data: dict[str, Any], message: str, session_key: str, correlation_id: str, state: AgentState
|
||||
self, response_data: dict[str, Any], message: str, thread_id: str, correlation_id: str, state: AgentState
|
||||
) -> dict[str, Any]:
|
||||
"""Build the success result returned to the HTTP caller."""
|
||||
return {
|
||||
"response": response_data.get("content"),
|
||||
"message": message,
|
||||
SESSION_ID_FIELD: session_key,
|
||||
THREAD_ID_FIELD: thread_id,
|
||||
"status": "success",
|
||||
"message_count": response_data.get("message_count", state.message_count),
|
||||
"correlationId": correlation_id,
|
||||
"correlation_id": correlation_id,
|
||||
}
|
||||
|
||||
def _build_request_data(
|
||||
self, req_body: dict[str, Any], message: str, conversation_id: str, correlation_id: str
|
||||
self, req_body: dict[str, Any], message: str, thread_id: str, correlation_id: str
|
||||
) -> dict[str, Any]:
|
||||
"""Create the durable entity request payload."""
|
||||
enable_tool_calls_value = req_body.get("enable_tool_calls")
|
||||
enable_tool_calls = True if enable_tool_calls_value is None else self._coerce_to_bool(enable_tool_calls_value)
|
||||
|
||||
role = self._coerce_chat_role(req_body.get("role"))
|
||||
|
||||
return RunRequest(
|
||||
message=message,
|
||||
role=role,
|
||||
role=req_body.get("role"),
|
||||
response_format=req_body.get("response_format"),
|
||||
enable_tool_calls=enable_tool_calls,
|
||||
conversation_id=conversation_id,
|
||||
thread_id=thread_id,
|
||||
correlation_id=correlation_id,
|
||||
).to_dict()
|
||||
|
||||
def _build_accepted_response(self, message: str, session_key: str, correlation_id: str) -> dict[str, Any]:
|
||||
def _build_accepted_response(self, message: str, thread_id: str, correlation_id: str) -> dict[str, Any]:
|
||||
"""Build the response returned when not waiting for completion."""
|
||||
return {
|
||||
"response": "Agent request accepted",
|
||||
"message": message,
|
||||
SESSION_ID_FIELD: session_key,
|
||||
THREAD_ID_FIELD: thread_id,
|
||||
"status": "accepted",
|
||||
"correlationId": correlation_id,
|
||||
"correlation_id": correlation_id,
|
||||
}
|
||||
|
||||
def _create_http_response(
|
||||
self,
|
||||
payload: dict[str, Any] | str,
|
||||
status_code: int,
|
||||
response_format: str,
|
||||
thread_id: str | None,
|
||||
) -> func.HttpResponse:
|
||||
"""Create the HTTP response using helper serializers for clarity."""
|
||||
if response_format == RESPONSE_FORMAT_TEXT:
|
||||
return self._build_plain_text_response(payload=payload, status_code=status_code, thread_id=thread_id)
|
||||
|
||||
return self._build_json_response(payload=payload, status_code=status_code)
|
||||
|
||||
def _build_plain_text_response(
|
||||
self,
|
||||
payload: dict[str, Any] | str,
|
||||
status_code: int,
|
||||
thread_id: str | None,
|
||||
) -> func.HttpResponse:
|
||||
"""Return a plain-text response with optional thread identifier header."""
|
||||
body_text = payload if isinstance(payload, str) else self._convert_payload_to_text(payload)
|
||||
headers = {"x-ms-thread-id": thread_id} if thread_id is not None else None
|
||||
return func.HttpResponse(body_text, status_code=status_code, mimetype="text/plain", headers=headers)
|
||||
|
||||
def _build_json_response(self, payload: dict[str, Any] | str, status_code: int) -> func.HttpResponse:
|
||||
"""Return the JSON response, serializing dictionaries as needed."""
|
||||
body_json = payload if isinstance(payload, str) else json.dumps(payload)
|
||||
return func.HttpResponse(body_json, status_code=status_code, mimetype="application/json")
|
||||
|
||||
def _convert_payload_to_text(self, payload: dict[str, Any]) -> str:
|
||||
"""Convert a structured payload into a human-readable text response."""
|
||||
for key in ("response", "error", "message"):
|
||||
value = payload.get(key)
|
||||
if isinstance(value, str) and value:
|
||||
return value
|
||||
return json.dumps(payload)
|
||||
|
||||
def _generate_unique_id(self) -> str:
|
||||
"""Generate a new unique identifier."""
|
||||
import uuid
|
||||
|
||||
return uuid.uuid4().hex
|
||||
|
||||
def _create_session_id(self, func_name: str, session_key: str | None) -> AgentSessionId:
|
||||
"""Create a session identifier using the provided key or a random value."""
|
||||
if session_key:
|
||||
return AgentSessionId(name=func_name, key=session_key)
|
||||
def _create_session_id(self, func_name: str, thread_id: str | None) -> AgentSessionId:
|
||||
"""Create a session identifier using the provided thread id or a random value."""
|
||||
if thread_id:
|
||||
return AgentSessionId(name=func_name, key=thread_id)
|
||||
return AgentSessionId.with_random_key(name=func_name)
|
||||
|
||||
def _resolve_session_key(self, req: func.HttpRequest, req_body: dict[str, Any]) -> str:
|
||||
"""Retrieve the session key from request body or query parameters."""
|
||||
def _resolve_thread_id(self, req: func.HttpRequest, req_body: dict[str, Any]) -> str:
|
||||
"""Retrieve the thread identifier from request body or query parameters."""
|
||||
params = req.params or {}
|
||||
|
||||
for key in SESSION_IDENTIFIER_KEYS:
|
||||
if key in req_body:
|
||||
value = req_body.get(key)
|
||||
if value is not None:
|
||||
return str(value)
|
||||
if THREAD_ID_FIELD in req_body:
|
||||
value = req_body.get(THREAD_ID_FIELD)
|
||||
if value is not None:
|
||||
return str(value)
|
||||
|
||||
for key in SESSION_IDENTIFIER_KEYS:
|
||||
if key in params:
|
||||
value = params.get(key)
|
||||
if value is not None:
|
||||
return str(value)
|
||||
if THREAD_ID_FIELD in params:
|
||||
value = params.get(THREAD_ID_FIELD)
|
||||
if value is not None:
|
||||
return str(value)
|
||||
|
||||
logger.debug("[HTTP Trigger] No session identifier provided; using random session key")
|
||||
logger.debug("[HTTP Trigger] No thread identifier provided; using random thread id")
|
||||
return self._generate_unique_id()
|
||||
|
||||
def _parse_incoming_request(self, req: func.HttpRequest) -> tuple[dict[str, Any], Any]:
|
||||
def _parse_incoming_request(self, req: func.HttpRequest) -> tuple[dict[str, Any], str, str]:
|
||||
"""Parse the incoming run request supporting JSON and plain text bodies."""
|
||||
headers = self._extract_normalized_headers(req)
|
||||
|
||||
normalized_content_type = self._extract_content_type(headers)
|
||||
body_parser, body_format = self._select_body_parser(normalized_content_type)
|
||||
prefers_json = self._accepts_json_response(headers)
|
||||
response_format = self._select_response_format(body_format=body_format, prefers_json=prefers_json)
|
||||
|
||||
req_body, message = body_parser(req)
|
||||
return req_body, message, response_format
|
||||
|
||||
def _extract_normalized_headers(self, req: func.HttpRequest) -> dict[str, str]:
|
||||
"""Create a lowercase header mapping from the incoming request."""
|
||||
headers: dict[str, str] = {}
|
||||
raw_headers = req.headers
|
||||
if isinstance(raw_headers, Mapping):
|
||||
headers_mapping = cast(Mapping[Any, Any], raw_headers)
|
||||
for key, value in headers_mapping.items():
|
||||
if value is not None:
|
||||
headers[str(key)] = str(value)
|
||||
|
||||
content_type_header = headers.get("content-type")
|
||||
|
||||
normalized_content_type = ""
|
||||
if content_type_header:
|
||||
normalized_content_type = content_type_header.split(";")[0].strip().lower()
|
||||
|
||||
if normalized_content_type in {"application/json"} or normalized_content_type.endswith("+json"):
|
||||
parser = self._parse_json_body
|
||||
else:
|
||||
parser = self._parse_text_body
|
||||
|
||||
return parser(req)
|
||||
headers[str(key).lower()] = str(value)
|
||||
return headers
|
||||
|
||||
@staticmethod
|
||||
def _parse_json_body(req: func.HttpRequest) -> tuple[dict[str, Any], Any]:
|
||||
def _extract_content_type(headers: dict[str, str]) -> str:
|
||||
"""Return the normalized content-type value (without parameters)."""
|
||||
content_type_header = headers.get("content-type", "")
|
||||
return content_type_header.split(";")[0].strip().lower() if content_type_header else ""
|
||||
|
||||
def _select_body_parser(
|
||||
self,
|
||||
normalized_content_type: str,
|
||||
) -> tuple[Callable[[func.HttpRequest], tuple[dict[str, Any], str]], str]:
|
||||
"""Choose the body parser and declared body format."""
|
||||
if normalized_content_type in {"application/json"} or normalized_content_type.endswith("+json"):
|
||||
return self._parse_json_body, RESPONSE_FORMAT_JSON
|
||||
return self._parse_text_body, RESPONSE_FORMAT_TEXT
|
||||
|
||||
@staticmethod
|
||||
def _accepts_json_response(headers: dict[str, str]) -> bool:
|
||||
"""Check whether the caller explicitly requests a JSON response."""
|
||||
accept_header = headers.get("accept")
|
||||
if not accept_header:
|
||||
return False
|
||||
|
||||
for value in accept_header.split(","):
|
||||
media_type = value.split(";")[0].strip().lower()
|
||||
if media_type == "application/json":
|
||||
return True
|
||||
return False
|
||||
|
||||
@staticmethod
|
||||
def _select_response_format(body_format: str, prefers_json: bool) -> str:
|
||||
"""Combine body format and accept preference to determine response format."""
|
||||
if body_format == RESPONSE_FORMAT_JSON or prefers_json:
|
||||
return RESPONSE_FORMAT_JSON
|
||||
return RESPONSE_FORMAT_TEXT
|
||||
|
||||
@staticmethod
|
||||
def _parse_json_body(req: func.HttpRequest) -> tuple[dict[str, Any], str]:
|
||||
req_body = req.get_json()
|
||||
if not isinstance(req_body, dict):
|
||||
raise IncomingRequestError("Invalid JSON payload. Expected an object.")
|
||||
@@ -603,46 +696,35 @@ class AgentFunctionApp(df.DFApp):
|
||||
return typed_req_body, message
|
||||
|
||||
@staticmethod
|
||||
def _parse_text_body(req: func.HttpRequest) -> tuple[dict[str, Any], Any]:
|
||||
def _parse_text_body(req: func.HttpRequest) -> tuple[dict[str, Any], str]:
|
||||
body_bytes = req.get_body()
|
||||
text_body = body_bytes.decode("utf-8", errors="replace") if body_bytes else ""
|
||||
message = text_body.strip()
|
||||
|
||||
if not message:
|
||||
raise IncomingRequestError("Message is required")
|
||||
|
||||
return {}, message
|
||||
|
||||
def _should_wait_for_completion(self, req: func.HttpRequest, req_body: dict[str, Any]) -> bool:
|
||||
"""Determine whether the caller requested to wait for completion."""
|
||||
def _should_wait_for_response(self, req: func.HttpRequest, req_body: dict[str, Any]) -> bool:
|
||||
"""Determine whether the caller requested to wait for the response."""
|
||||
header_value = None
|
||||
raw_headers = req.headers
|
||||
if isinstance(raw_headers, Mapping):
|
||||
headers_mapping = cast(Mapping[Any, Any], raw_headers)
|
||||
for key, value in headers_mapping.items():
|
||||
if str(key).lower() == "x-wait-for-completion":
|
||||
if str(key).lower() == WAIT_FOR_RESPONSE_HEADER:
|
||||
header_value = value
|
||||
break
|
||||
|
||||
if header_value is not None:
|
||||
return self._coerce_to_bool(header_value)
|
||||
|
||||
for key in ("wait_for_completion", "waitForCompletion", "WaitForCompletion"):
|
||||
if key in req_body:
|
||||
return self._coerce_to_bool(req_body.get(key))
|
||||
params = req.params or {}
|
||||
if WAIT_FOR_RESPONSE_FIELD in params:
|
||||
return self._coerce_to_bool(params.get(WAIT_FOR_RESPONSE_FIELD))
|
||||
|
||||
return False
|
||||
if WAIT_FOR_RESPONSE_FIELD in req_body:
|
||||
return self._coerce_to_bool(req_body.get(WAIT_FOR_RESPONSE_FIELD))
|
||||
|
||||
def _coerce_chat_role(self, value: Any) -> ChatRole:
|
||||
"""Convert user-provided role to ChatRole, defaulting to user on error."""
|
||||
if isinstance(value, ChatRole):
|
||||
return value
|
||||
if isinstance(value, str):
|
||||
try:
|
||||
return ChatRole(value.strip().lower())
|
||||
except ValueError:
|
||||
logger.warning("[AgentFunctionApp] Invalid role '%s'; defaulting to user", value)
|
||||
return ChatRole.USER
|
||||
return True
|
||||
|
||||
def _coerce_to_bool(self, value: Any) -> bool:
|
||||
"""Convert various representations into a boolean flag."""
|
||||
|
||||
@@ -20,7 +20,7 @@ class AgentCallbackContext:
|
||||
|
||||
agent_name: str
|
||||
correlation_id: str
|
||||
conversation_id: str | None = None
|
||||
thread_id: str | None = None
|
||||
request_message: str | None = None
|
||||
|
||||
|
||||
|
||||
@@ -14,10 +14,10 @@ from collections.abc import AsyncIterable
|
||||
from typing import Any, cast
|
||||
|
||||
import azure.durable_functions as df
|
||||
from agent_framework import AgentProtocol, AgentRunResponse, AgentRunResponseUpdate, get_logger
|
||||
from agent_framework import AgentProtocol, AgentRunResponse, AgentRunResponseUpdate, Role, get_logger
|
||||
|
||||
from ._callbacks import AgentCallbackContext, AgentResponseCallbackProtocol
|
||||
from ._models import AgentResponse, ChatRole, RunRequest
|
||||
from ._models import AgentResponse, RunRequest
|
||||
from ._state import AgentState
|
||||
|
||||
logger = get_logger("agent_framework.azurefunctions.entities")
|
||||
@@ -81,33 +81,32 @@ class AgentEntity:
|
||||
"""
|
||||
# Convert string or dict to RunRequest
|
||||
if isinstance(request, str):
|
||||
run_request = RunRequest(message=request, role=ChatRole.USER)
|
||||
run_request = RunRequest(message=request, role=Role.USER)
|
||||
elif isinstance(request, dict):
|
||||
run_request = RunRequest.from_dict(request)
|
||||
else:
|
||||
run_request = request
|
||||
|
||||
message = run_request.message
|
||||
conversation_id = run_request.conversation_id
|
||||
thread_id = run_request.thread_id
|
||||
correlation_id = run_request.correlation_id
|
||||
if not conversation_id:
|
||||
raise ValueError("RunRequest must include a conversation_id")
|
||||
if not thread_id:
|
||||
raise ValueError("RunRequest must include a thread_id")
|
||||
if not correlation_id:
|
||||
raise ValueError("RunRequest must include a correlation_id")
|
||||
role = run_request.role or ChatRole.USER
|
||||
role = run_request.role or Role.USER
|
||||
response_format = run_request.response_format
|
||||
enable_tool_calls = run_request.enable_tool_calls
|
||||
|
||||
logger.debug(f"[AgentEntity.run_agent] Received message: {message}")
|
||||
logger.debug(f"[AgentEntity.run_agent] Conversation ID: {conversation_id}")
|
||||
logger.debug(f"[AgentEntity.run_agent] Thread ID: {thread_id}")
|
||||
logger.debug(f"[AgentEntity.run_agent] Correlation ID: {correlation_id}")
|
||||
logger.debug(f"[AgentEntity.run_agent] Role: {role.value if isinstance(role, ChatRole) else role}")
|
||||
logger.debug(f"[AgentEntity.run_agent] Role: {role.value}")
|
||||
logger.debug(f"[AgentEntity.run_agent] Enable tool calls: {enable_tool_calls}")
|
||||
logger.debug(f"[AgentEntity.run_agent] Response format: {'provided' if response_format else 'none'}")
|
||||
|
||||
# Store message in history with role
|
||||
role_str = role.value if isinstance(role, ChatRole) else role
|
||||
self.state.add_user_message(message, role=role_str, correlation_id=correlation_id)
|
||||
self.state.add_user_message(message, role=role, correlation_id=correlation_id)
|
||||
|
||||
logger.debug("[AgentEntity.run_agent] Executing agent...")
|
||||
|
||||
@@ -123,7 +122,7 @@ class AgentEntity:
|
||||
agent_run_response: AgentRunResponse = await self._invoke_agent(
|
||||
run_kwargs=run_kwargs,
|
||||
correlation_id=correlation_id,
|
||||
conversation_id=conversation_id,
|
||||
thread_id=thread_id,
|
||||
request_message=message,
|
||||
)
|
||||
|
||||
@@ -160,7 +159,7 @@ class AgentEntity:
|
||||
agent_response = AgentResponse(
|
||||
response=response_text,
|
||||
message=str(message),
|
||||
conversation_id=str(conversation_id),
|
||||
thread_id=str(thread_id),
|
||||
status="success",
|
||||
message_count=self.state.message_count,
|
||||
structured_response=structured_response,
|
||||
@@ -185,7 +184,7 @@ class AgentEntity:
|
||||
error_response = AgentResponse(
|
||||
response=f"Error: {exc!s}",
|
||||
message=str(message),
|
||||
conversation_id=str(conversation_id),
|
||||
thread_id=str(thread_id),
|
||||
status="error",
|
||||
message_count=self.state.message_count,
|
||||
error=str(exc),
|
||||
@@ -197,7 +196,7 @@ class AgentEntity:
|
||||
self,
|
||||
run_kwargs: dict[str, Any],
|
||||
correlation_id: str,
|
||||
conversation_id: str,
|
||||
thread_id: str,
|
||||
request_message: str,
|
||||
) -> AgentRunResponse:
|
||||
"""Execute the agent, preferring streaming when available."""
|
||||
@@ -205,7 +204,7 @@ class AgentEntity:
|
||||
if self.callback is not None:
|
||||
callback_context = self._build_callback_context(
|
||||
correlation_id=correlation_id,
|
||||
conversation_id=conversation_id,
|
||||
thread_id=thread_id,
|
||||
request_message=request_message,
|
||||
)
|
||||
|
||||
@@ -319,7 +318,7 @@ class AgentEntity:
|
||||
def _build_callback_context(
|
||||
self,
|
||||
correlation_id: str,
|
||||
conversation_id: str,
|
||||
thread_id: str,
|
||||
request_message: str,
|
||||
) -> AgentCallbackContext:
|
||||
"""Create the callback context provided to consumers."""
|
||||
@@ -327,7 +326,7 @@ class AgentEntity:
|
||||
return AgentCallbackContext(
|
||||
agent_name=agent_name,
|
||||
correlation_id=correlation_id,
|
||||
conversation_id=conversation_id,
|
||||
thread_id=thread_id,
|
||||
request_message=request_message,
|
||||
)
|
||||
|
||||
@@ -375,11 +374,8 @@ def create_agent_entity(
|
||||
if operation == "run_agent":
|
||||
input_data: Any = context.get_input()
|
||||
|
||||
# Support both old format (message + conversation_id) and new format (RunRequest dict)
|
||||
# This provides backward compatibility
|
||||
request: str | dict[str, Any]
|
||||
if isinstance(input_data, dict) and "message" in input_data:
|
||||
# Input can be either old format or new RunRequest format
|
||||
request = cast(dict[str, Any], input_data)
|
||||
else:
|
||||
# Fall back to treating input as message string
|
||||
|
||||
@@ -5,21 +5,23 @@
|
||||
This module defines the request and response models used by the framework.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import inspect
|
||||
import uuid
|
||||
from collections.abc import MutableMapping
|
||||
from dataclasses import dataclass
|
||||
from enum import Enum
|
||||
from importlib import import_module
|
||||
from typing import TYPE_CHECKING, Any, cast
|
||||
|
||||
import azure.durable_functions as df
|
||||
from agent_framework import AgentThread
|
||||
from agent_framework import AgentThread, Role
|
||||
|
||||
if TYPE_CHECKING: # pragma: no cover - type checking imports only
|
||||
from pydantic import BaseModel
|
||||
|
||||
_PydanticBaseModel: type["BaseModel"] | None
|
||||
_PydanticBaseModel: type[BaseModel] | None
|
||||
|
||||
try:
|
||||
from pydantic import BaseModel as _RuntimeBaseModel
|
||||
except ImportError: # pragma: no cover - optional dependency
|
||||
@@ -28,14 +30,6 @@ else:
|
||||
_PydanticBaseModel = _RuntimeBaseModel
|
||||
|
||||
|
||||
class ChatRole(str, Enum):
|
||||
"""Chat message role enum."""
|
||||
|
||||
USER = "user"
|
||||
SYSTEM = "system"
|
||||
ASSISTANT = "assistant"
|
||||
|
||||
|
||||
@dataclass
|
||||
class AgentSessionId:
|
||||
"""Represents an agent session ID, which is used to identify a long-running agent session.
|
||||
@@ -63,7 +57,7 @@ class AgentSessionId:
|
||||
return f"{AgentSessionId.ENTITY_NAME_PREFIX}{name}"
|
||||
|
||||
@staticmethod
|
||||
def with_random_key(name: str) -> "AgentSessionId":
|
||||
def with_random_key(name: str) -> AgentSessionId:
|
||||
"""Creates a new AgentSessionId with the specified name and a randomly generated key.
|
||||
|
||||
Args:
|
||||
@@ -83,7 +77,7 @@ class AgentSessionId:
|
||||
return df.EntityId(self.to_entity_name(self.name), self.key)
|
||||
|
||||
@staticmethod
|
||||
def from_entity_id(entity_id: df.EntityId) -> "AgentSessionId":
|
||||
def from_entity_id(entity_id: df.EntityId) -> AgentSessionId:
|
||||
"""Creates an AgentSessionId from a Durable Functions EntityId.
|
||||
|
||||
Args:
|
||||
@@ -113,7 +107,7 @@ class AgentSessionId:
|
||||
return f"AgentSessionId(name='{self.name}', key='{self.key}')"
|
||||
|
||||
@staticmethod
|
||||
def parse(session_id_string: str) -> "AgentSessionId":
|
||||
def parse(session_id_string: str) -> AgentSessionId:
|
||||
"""Parses a string representation of an agent session ID.
|
||||
|
||||
Args:
|
||||
@@ -172,7 +166,7 @@ class DurableAgentThread(AgentThread):
|
||||
service_thread_id: str | None = None,
|
||||
message_store: Any = None,
|
||||
context_provider: Any = None,
|
||||
) -> "DurableAgentThread":
|
||||
) -> DurableAgentThread:
|
||||
"""Creates a durable thread pre-associated with the supplied session ID."""
|
||||
return cls(
|
||||
session_id=session_id,
|
||||
@@ -195,7 +189,7 @@ class DurableAgentThread(AgentThread):
|
||||
*,
|
||||
message_store: Any = None,
|
||||
**kwargs: Any,
|
||||
) -> "DurableAgentThread":
|
||||
) -> DurableAgentThread:
|
||||
"""Restores a durable thread, rehydrating the stored session identifier."""
|
||||
state_payload = dict(serialized_thread_state)
|
||||
session_id_value = state_payload.pop(cls._SERIALIZED_SESSION_ID_KEY, None)
|
||||
@@ -217,7 +211,7 @@ class DurableAgentThread(AgentThread):
|
||||
return thread
|
||||
|
||||
|
||||
def _serialize_response_format(response_format: type["BaseModel"] | None) -> Any:
|
||||
def _serialize_response_format(response_format: type[BaseModel] | None) -> Any:
|
||||
"""Serialize response format for transport across durable function boundaries."""
|
||||
if response_format is None:
|
||||
return None
|
||||
@@ -235,7 +229,7 @@ def _serialize_response_format(response_format: type["BaseModel"] | None) -> Any
|
||||
}
|
||||
|
||||
|
||||
def _deserialize_response_format(response_format: Any) -> type["BaseModel"] | None:
|
||||
def _deserialize_response_format(response_format: Any) -> type[BaseModel] | None:
|
||||
"""Deserialize response format back into actionable type if possible."""
|
||||
if response_format is None:
|
||||
return None
|
||||
@@ -287,17 +281,45 @@ class RunRequest:
|
||||
role: The role of the message sender (user, system, or assistant)
|
||||
response_format: Optional Pydantic BaseModel type describing the structured response format
|
||||
enable_tool_calls: Whether to enable tool calls for this request
|
||||
conversation_id: Optional conversation/session ID for tracking
|
||||
thread_id: Optional thread ID for tracking
|
||||
correlation_id: Optional correlation ID for tracking the response to this specific request
|
||||
"""
|
||||
|
||||
message: str
|
||||
role: ChatRole = ChatRole.USER
|
||||
response_format: type["BaseModel"] | None = None
|
||||
role: Role = Role.USER
|
||||
response_format: type[BaseModel] | None = None
|
||||
enable_tool_calls: bool = True
|
||||
conversation_id: str | None = None
|
||||
thread_id: str | None = None
|
||||
correlation_id: str | None = None
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
message: str,
|
||||
role: Role | str | None = Role.USER,
|
||||
response_format: type[BaseModel] | None = None,
|
||||
enable_tool_calls: bool = True,
|
||||
thread_id: str | None = None,
|
||||
correlation_id: str | None = None,
|
||||
) -> None:
|
||||
self.message = message
|
||||
self.role = self.coerce_role(role)
|
||||
self.response_format = response_format
|
||||
self.enable_tool_calls = enable_tool_calls
|
||||
self.thread_id = thread_id
|
||||
self.correlation_id = correlation_id
|
||||
|
||||
@staticmethod
|
||||
def coerce_role(value: Role | str | None) -> Role:
|
||||
"""Normalize various role representations into a Role instance."""
|
||||
if isinstance(value, Role):
|
||||
return value
|
||||
if isinstance(value, str):
|
||||
normalized = value.strip()
|
||||
if not normalized:
|
||||
return Role.USER
|
||||
return Role(value=normalized.lower())
|
||||
return Role.USER
|
||||
|
||||
def to_dict(self) -> dict[str, Any]:
|
||||
"""Convert to dictionary for JSON serialization."""
|
||||
result = {
|
||||
@@ -307,30 +329,21 @@ class RunRequest:
|
||||
}
|
||||
if self.response_format:
|
||||
result["response_format"] = _serialize_response_format(self.response_format)
|
||||
if self.conversation_id:
|
||||
result["conversation_id"] = self.conversation_id
|
||||
if self.thread_id:
|
||||
result["thread_id"] = self.thread_id
|
||||
if self.correlation_id:
|
||||
result["correlation_id"] = self.correlation_id
|
||||
return result
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, data: dict[str, Any]) -> "RunRequest":
|
||||
def from_dict(cls, data: dict[str, Any]) -> RunRequest:
|
||||
"""Create RunRequest from dictionary."""
|
||||
role_str = data.get("role")
|
||||
if role_str:
|
||||
try:
|
||||
role = ChatRole(role_str.lower())
|
||||
except ValueError:
|
||||
role = ChatRole.USER # Default to USER if invalid
|
||||
else:
|
||||
role = ChatRole.USER
|
||||
|
||||
return cls(
|
||||
message=data.get("message", ""),
|
||||
role=role,
|
||||
role=cls.coerce_role(data.get("role")),
|
||||
response_format=_deserialize_response_format(data.get("response_format")),
|
||||
enable_tool_calls=data.get("enable_tool_calls", True),
|
||||
conversation_id=data.get("conversation_id"),
|
||||
thread_id=data.get("thread_id"),
|
||||
correlation_id=data.get("correlation_id"),
|
||||
)
|
||||
|
||||
@@ -342,7 +355,7 @@ class AgentResponse:
|
||||
Attributes:
|
||||
response: The agent's text response (or None for structured responses)
|
||||
message: The original message sent to the agent
|
||||
conversation_id: The conversation/session ID
|
||||
thread_id: The thread identifier
|
||||
status: Status of the execution (success, error, etc.)
|
||||
message_count: Number of messages in the conversation
|
||||
error: Error message if status is error
|
||||
@@ -352,7 +365,7 @@ class AgentResponse:
|
||||
|
||||
response: str | None
|
||||
message: str
|
||||
conversation_id: str | None
|
||||
thread_id: str | None
|
||||
status: str
|
||||
message_count: int = 0
|
||||
error: str | None = None
|
||||
@@ -363,7 +376,7 @@ class AgentResponse:
|
||||
"""Convert to dictionary for JSON serialization."""
|
||||
result = {
|
||||
"message": self.message,
|
||||
"conversation_id": self.conversation_id,
|
||||
"thread_id": self.thread_id,
|
||||
"status": self.status,
|
||||
"message_count": self.message_count,
|
||||
}
|
||||
|
||||
@@ -136,7 +136,7 @@ class DurableAIAgent(AgentProtocol):
|
||||
message=message_str,
|
||||
enable_tool_calls=enable_tool_calls,
|
||||
correlation_id=correlation_id,
|
||||
conversation_id=session_id.key,
|
||||
thread_id=session_id.key,
|
||||
response_format=response_format,
|
||||
)
|
||||
|
||||
|
||||
@@ -8,9 +8,9 @@ serializing agent framework responses.
|
||||
|
||||
from collections.abc import MutableMapping
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any, Literal, cast
|
||||
from typing import Any, cast
|
||||
|
||||
from agent_framework import AgentRunResponse, ChatMessage, get_logger
|
||||
from agent_framework import AgentRunResponse, ChatMessage, Role, get_logger
|
||||
|
||||
logger = get_logger("agent_framework.azurefunctions.state")
|
||||
|
||||
@@ -38,7 +38,7 @@ class AgentState:
|
||||
def add_user_message(
|
||||
self,
|
||||
content: str,
|
||||
role: Literal["user", "system", "assistant", "tool"] = "user",
|
||||
role: Role = Role.USER,
|
||||
correlation_id: str | None = None,
|
||||
) -> None:
|
||||
"""Add a user message to the conversation history as a ChatMessage object.
|
||||
|
||||
Reference in New Issue
Block a user