mirror of
https://github.com/microsoft/agent-framework.git
synced 2026-06-16 21:04:09 +08:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
587356c778 | ||
|
|
1b7940c91e | ||
|
|
2f4c4aa614 | ||
|
|
052ba7be07 | ||
|
|
c67d3523ae | ||
|
|
83ce6a9602 | ||
|
|
50fdcbaf57 | ||
|
|
67b0282813 | ||
|
|
0009e330af | ||
|
|
a4b9539b62 | ||
|
|
b7990908fe | ||
|
|
84bae0f42a |
@@ -0,0 +1,216 @@
|
||||
# Probe the highest allowed dependency versions, then open issues/PRs from the passing updates.
|
||||
name: Python - Dependency Range Validation
|
||||
|
||||
on:
|
||||
workflow_dispatch:
|
||||
|
||||
permissions:
|
||||
contents: write
|
||||
issues: write
|
||||
pull-requests: write
|
||||
|
||||
env:
|
||||
UV_CACHE_DIR: /tmp/.uv-cache
|
||||
|
||||
jobs:
|
||||
dependency-range-validation:
|
||||
name: Dependency Range Validation
|
||||
runs-on: ubuntu-latest
|
||||
env:
|
||||
# For now only run 3.13, if we do encounter situations where there are mismatches between packages and python versions (other then 3.10 and 3.14 which are known to not be able to install everything)
|
||||
# then we will have to reevaluate.
|
||||
UV_PYTHON: "3.13"
|
||||
GH_TOKEN: ${{ secrets.GITHUB_TOKEN }}
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
with:
|
||||
fetch-depth: 0
|
||||
|
||||
- name: Set up python and install the project
|
||||
uses: ./.github/actions/python-setup
|
||||
with:
|
||||
python-version: ${{ env.UV_PYTHON }}
|
||||
os: ${{ runner.os }}
|
||||
env:
|
||||
UV_CACHE_DIR: /tmp/.uv-cache
|
||||
|
||||
- name: Run dependency range validation
|
||||
id: validate_ranges
|
||||
# Keep workflow running so we can still publish diagnostics from this run.
|
||||
continue-on-error: true
|
||||
run: uv run poe validate-dependency-bounds-project --mode upper --project "*"
|
||||
working-directory: ./python
|
||||
|
||||
- name: Upload dependency range report
|
||||
# Always publish the report so failures are inspectable even when validation fails.
|
||||
if: always()
|
||||
uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: dependency-range-results
|
||||
path: python/scripts/dependencies/dependency-range-results.json
|
||||
if-no-files-found: warn
|
||||
|
||||
- name: Create issues for failed dependency candidates
|
||||
# Always process the report so failed candidates create actionable tracking issues.
|
||||
if: always()
|
||||
uses: actions/github-script@v8
|
||||
with:
|
||||
script: |
|
||||
const fs = require("fs")
|
||||
const reportPath = "python/scripts/dependencies/dependency-range-results.json"
|
||||
|
||||
if (!fs.existsSync(reportPath)) {
|
||||
core.warning(`No dependency range report found at ${reportPath}`)
|
||||
return
|
||||
}
|
||||
|
||||
const report = JSON.parse(fs.readFileSync(reportPath, "utf8"))
|
||||
const dependencyFailures = []
|
||||
|
||||
for (const packageResult of report.packages ?? []) {
|
||||
for (const dependency of packageResult.dependencies ?? []) {
|
||||
const candidateVersions = new Set(dependency.candidate_versions ?? [])
|
||||
const failedAttempts = (dependency.attempts ?? []).filter(
|
||||
(attempt) => attempt.status === "failed" && candidateVersions.has(attempt.trial_upper)
|
||||
)
|
||||
if (!failedAttempts.length) {
|
||||
continue
|
||||
}
|
||||
|
||||
const failuresByVersion = new Map()
|
||||
for (const attempt of failedAttempts) {
|
||||
const version = attempt.trial_upper || "unknown"
|
||||
if (!failuresByVersion.has(version)) {
|
||||
failuresByVersion.set(version, attempt.error || "No error output captured.")
|
||||
}
|
||||
}
|
||||
|
||||
dependencyFailures.push({
|
||||
packageName: packageResult.package_name,
|
||||
projectPath: packageResult.project_path,
|
||||
dependencyName: dependency.name,
|
||||
originalRequirements: dependency.original_requirements ?? [],
|
||||
finalRequirements: dependency.final_requirements ?? [],
|
||||
failedVersions: [...failuresByVersion.entries()].map(([version, error]) => ({ version, error })),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
if (!dependencyFailures.length) {
|
||||
core.info("No failing dependency candidates found.")
|
||||
return
|
||||
}
|
||||
|
||||
const owner = context.repo.owner
|
||||
const repo = context.repo.repo
|
||||
const openIssues = await github.paginate(github.rest.issues.listForRepo, {
|
||||
owner,
|
||||
repo,
|
||||
state: "open",
|
||||
per_page: 100,
|
||||
})
|
||||
const openIssueTitles = new Set(
|
||||
openIssues.filter((issue) => !issue.pull_request).map((issue) => issue.title)
|
||||
)
|
||||
|
||||
const formatError = (message) => String(message || "No error output captured.").replace(/```/g, "'''")
|
||||
|
||||
for (const failure of dependencyFailures) {
|
||||
const title = `Dependency validation failed: ${failure.dependencyName} (${failure.packageName})`
|
||||
if (openIssueTitles.has(title)) {
|
||||
core.info(`Issue already exists: ${title}`)
|
||||
continue
|
||||
}
|
||||
|
||||
const visibleFailures = failure.failedVersions.slice(0, 5)
|
||||
const omittedCount = failure.failedVersions.length - visibleFailures.length
|
||||
const failureDetails = visibleFailures
|
||||
.map(
|
||||
(entry) =>
|
||||
`- \`${entry.version}\`\n\n\`\`\`\n${formatError(entry.error).slice(0, 3500)}\n\`\`\``
|
||||
)
|
||||
.join("\n\n")
|
||||
|
||||
const body = [
|
||||
"Automated dependency range validation found candidate versions that failed checks.",
|
||||
"",
|
||||
`- Package: \`${failure.packageName}\``,
|
||||
`- Project path: \`${failure.projectPath}\``,
|
||||
`- Dependency: \`${failure.dependencyName}\``,
|
||||
`- Original requirements: ${
|
||||
failure.originalRequirements.length
|
||||
? failure.originalRequirements.map((value) => `\`${value}\``).join(", ")
|
||||
: "_none_"
|
||||
}`,
|
||||
`- Final requirements after run: ${
|
||||
failure.finalRequirements.length
|
||||
? failure.finalRequirements.map((value) => `\`${value}\``).join(", ")
|
||||
: "_none_"
|
||||
}`,
|
||||
"",
|
||||
"### Failed versions and errors",
|
||||
failureDetails,
|
||||
omittedCount > 0 ? `\n_Additional failed versions omitted: ${omittedCount}_` : "",
|
||||
"",
|
||||
`Workflow run: ${context.serverUrl}/${owner}/${repo}/actions/runs/${context.runId}`,
|
||||
].join("\n")
|
||||
|
||||
await github.rest.issues.create({
|
||||
owner,
|
||||
repo,
|
||||
title,
|
||||
body,
|
||||
})
|
||||
openIssueTitles.add(title)
|
||||
core.info(`Created issue: ${title}`)
|
||||
}
|
||||
|
||||
- name: Refresh lockfile
|
||||
# Only refresh lockfile after a clean validation to avoid committing known-bad ranges.
|
||||
if: steps.validate_ranges.outcome == 'success'
|
||||
run: uv lock --upgrade
|
||||
working-directory: ./python
|
||||
|
||||
- name: Commit and push dependency updates
|
||||
id: commit_updates
|
||||
if: steps.validate_ranges.outcome == 'success'
|
||||
run: |
|
||||
BRANCH="automation/python-dependency-range-updates"
|
||||
|
||||
git config user.name "github-actions[bot]"
|
||||
git config user.email "41898282+github-actions[bot]@users.noreply.github.com"
|
||||
git checkout -B "${BRANCH}"
|
||||
|
||||
git add python/packages/*/pyproject.toml python/uv.lock
|
||||
if git diff --cached --quiet; then
|
||||
echo "has_changes=false" >> "$GITHUB_OUTPUT"
|
||||
echo "No dependency updates to commit."
|
||||
exit 0
|
||||
fi
|
||||
|
||||
git commit -m "chore: update dependency ranges"
|
||||
git push --force-with-lease --set-upstream origin "${BRANCH}"
|
||||
echo "has_changes=true" >> "$GITHUB_OUTPUT"
|
||||
|
||||
- name: Create or update pull request with GitHub CLI
|
||||
# Only open/update PRs for validated updates to keep automation branches trustworthy.
|
||||
if: steps.validate_ranges.outcome == 'success' && steps.commit_updates.outputs.has_changes == 'true'
|
||||
run: |
|
||||
BRANCH="automation/python-dependency-range-updates"
|
||||
PR_TITLE="Python: chore: update dependency ranges"
|
||||
PR_BODY_FILE="$(mktemp)"
|
||||
|
||||
cat > "${PR_BODY_FILE}" <<'EOF'
|
||||
This PR was generated by the dependency range validation workflow.
|
||||
|
||||
- Ran `uv run poe validate-dependency-bounds-project --mode upper --project "*"`
|
||||
- Updated package dependency bounds
|
||||
- Refreshed `python/uv.lock` with `uv lock --upgrade`
|
||||
EOF
|
||||
|
||||
PR_NUMBER="$(gh pr list --head "${BRANCH}" --base main --state open --json number --jq '.[0].number')"
|
||||
if [ -n "${PR_NUMBER}" ]; then
|
||||
gh pr edit "${PR_NUMBER}" --title "${PR_TITLE}" --body-file "${PR_BODY_FILE}"
|
||||
else
|
||||
gh pr create --base main --head "${BRANCH}" --title "${PR_TITLE}" --body-file "${PR_BODY_FILE}"
|
||||
fi
|
||||
@@ -0,0 +1,91 @@
|
||||
name: Python - Dev Dependency Upgrade
|
||||
|
||||
on:
|
||||
workflow_dispatch:
|
||||
|
||||
permissions:
|
||||
contents: write
|
||||
pull-requests: write
|
||||
|
||||
env:
|
||||
UV_CACHE_DIR: /tmp/.uv-cache
|
||||
|
||||
jobs:
|
||||
upgrade-dev-dependencies:
|
||||
name: Upgrade Dev Dependencies
|
||||
runs-on: ubuntu-latest
|
||||
env:
|
||||
UV_PYTHON: "3.13"
|
||||
GH_TOKEN: ${{ secrets.GITHUB_TOKEN }}
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
with:
|
||||
fetch-depth: 0
|
||||
|
||||
- name: Set up python and install the project
|
||||
uses: ./.github/actions/python-setup
|
||||
with:
|
||||
python-version: ${{ env.UV_PYTHON }}
|
||||
os: ${{ runner.os }}
|
||||
env:
|
||||
UV_CACHE_DIR: /tmp/.uv-cache
|
||||
|
||||
- name: Upgrade dev dependencies and validate workspace
|
||||
run: uv run poe upgrade-dev-dependencies
|
||||
working-directory: ./python
|
||||
|
||||
- name: Commit and push dev dependency updates
|
||||
id: commit_updates
|
||||
run: |
|
||||
BRANCH="automation/python-dev-dependency-updates"
|
||||
|
||||
git config user.name "github-actions[bot]"
|
||||
git config user.email "41898282+github-actions[bot]@users.noreply.github.com"
|
||||
git checkout -B "${BRANCH}"
|
||||
|
||||
git add python/pyproject.toml python/packages/*/pyproject.toml python/uv.lock
|
||||
if git diff --cached --quiet; then
|
||||
echo "has_changes=false" >> "$GITHUB_OUTPUT"
|
||||
echo "No dev dependency updates to commit."
|
||||
exit 0
|
||||
fi
|
||||
|
||||
git commit -F- <<'EOF'
|
||||
Python: chore: upgrade dev dependencies
|
||||
EOF
|
||||
git push --force-with-lease --set-upstream origin "${BRANCH}"
|
||||
echo "has_changes=true" >> "$GITHUB_OUTPUT"
|
||||
|
||||
- name: Create or update pull request with GitHub CLI
|
||||
if: steps.commit_updates.outputs.has_changes == 'true'
|
||||
run: |
|
||||
BRANCH="automation/python-dev-dependency-updates"
|
||||
PR_TITLE="Python: chore: upgrade dev dependencies"
|
||||
PR_BODY_FILE="$(mktemp)"
|
||||
|
||||
cat > "${PR_BODY_FILE}" <<'EOF'
|
||||
### Motivation and Context
|
||||
|
||||
This automated update refreshes Python dev dependency pins across the workspace and reruns the repo validation gates before opening a pull request.
|
||||
|
||||
### Description
|
||||
|
||||
- Ran `uv run poe upgrade-dev-dependencies`
|
||||
- Refreshed dev dependency pins in workspace `pyproject.toml` files
|
||||
- Refreshed `python/uv.lock` with `uv lock --upgrade`
|
||||
- Reinstalled from the frozen lockfile and reran `check`, `typing`, and `test`
|
||||
|
||||
### Contribution Checklist
|
||||
|
||||
- [x] The code builds clean without any errors or warnings
|
||||
- [x] The PR follows the [Contribution Guidelines](https://github.com/microsoft/agent-framework/blob/main/CONTRIBUTING.md)
|
||||
- [x] All unit tests pass, and I have added new tests where possible
|
||||
- [ ] **Is this a breaking change?** If yes, add "[BREAKING]" prefix to the title of the PR.
|
||||
EOF
|
||||
|
||||
PR_NUMBER="$(gh pr list --head "${BRANCH}" --base main --state open --json number --jq '.[0].number')"
|
||||
if [ -n "${PR_NUMBER}" ]; then
|
||||
gh pr edit "${PR_NUMBER}" --title "${PR_TITLE}" --body-file "${PR_BODY_FILE}"
|
||||
else
|
||||
gh pr create --base main --head "${BRANCH}" --title "${PR_TITLE}" --body-file "${PR_BODY_FILE}"
|
||||
fi
|
||||
@@ -76,6 +76,9 @@ jobs:
|
||||
- name: Run lab tests
|
||||
run: cd packages/lab && uv run poe test
|
||||
|
||||
- name: Run resource-intensive lab tests
|
||||
run: cd packages/lab && uv run pytest -m "resource_intensive and not integration" --junitxml=test-results-resource-intensive.xml
|
||||
|
||||
- name: Run lab lint
|
||||
run: cd packages/lab && uv run poe lint
|
||||
|
||||
|
||||
@@ -205,6 +205,9 @@ WARP.md
|
||||
**/memory-bank/
|
||||
**/projectBrief.md
|
||||
**/tmpclaude*
|
||||
# Dependency-bound validation reports
|
||||
python/scripts/dependency-*-results.json
|
||||
python/scripts/dependencies/dependency-*-results.json
|
||||
|
||||
# Azurite storage emulator files
|
||||
*/__azurite_db_blob__.json*
|
||||
|
||||
@@ -4,8 +4,8 @@ status: accepted
|
||||
contact: westey-m
|
||||
date: 2025-07-10 {YYYY-MM-DD when the decision was last updated}
|
||||
deciders: sergeymenshykh, markwallace, rbarreto, dmytrostruk, westey-m, eavanvalkenburg, stephentoub
|
||||
consulted:
|
||||
informed:
|
||||
consulted:
|
||||
informed:
|
||||
---
|
||||
|
||||
# Agent Run Responses Design
|
||||
@@ -64,7 +64,7 @@ Approaches observed from the compared SDKs:
|
||||
| AutoGen | **Approach 1** Separates messages into Agent-Agent (maps to Primary) and Internal (maps to Secondary) and these are returned as separate properties on the agent response object. See [types of messages](https://microsoft.github.io/autogen/stable/user-guide/agentchat-user-guide/tutorial/messages.html#types-of-messages) and [Response](https://microsoft.github.io/autogen/stable/reference/python/autogen_agentchat.base.html#autogen_agentchat.base.Response) | **Approach 2** Returns a stream of internal events and the last item is a Response object. See [ChatAgent.on_messages_stream](https://microsoft.github.io/autogen/stable/reference/python/autogen_agentchat.base.html#autogen_agentchat.base.ChatAgent.on_messages_stream) |
|
||||
| OpenAI Agent SDK | **Approach 1** Separates new_items (Primary+Secondary) from final output (Primary) as separate properties on the [RunResult](https://github.com/openai/openai-agents-python/blob/main/src/agents/result.py#L39) | **Approach 1** Similar to non-streaming, has a way of streaming updates via a method on the response object which includes all data, and then a separate final output property on the response object which is populated only when the run is complete. See [RunResultStreaming](https://github.com/openai/openai-agents-python/blob/main/src/agents/result.py#L136) |
|
||||
| Google ADK | **Approach 2** [Emits events](https://google.github.io/adk-docs/runtime/#step-by-step-breakdown) with [FinalResponse](https://github.com/google/adk-java/blob/main/core/src/main/java/com/google/adk/events/Event.java#L232) true (Primary) / false (Secondary) and callers have to filter out those with false to get just the final response message | **Approach 2** Similar to non-streaming except [events](https://google.github.io/adk-docs/runtime/#streaming-vs-non-streaming-output-partialtrue) are emitted with [Partial](https://github.com/google/adk-java/blob/main/core/src/main/java/com/google/adk/events/Event.java#L133) true to indicate that they are streaming messages. A final non partial event is also emitted. |
|
||||
| AWS (Strands) | **Approach 3** Returns an [AgentResult](https://strandsagents.com/docs/api/python/strands.agent.agent_result/#agentresult) (Primary) with messages and a reason for the run's completion. | **Approach 2** [Streams events](https://strandsagents.com/docs/user-guide/concepts/streaming/) (Primary+Secondary) including, response text, current_tool_use, even data from "callbacks" (strands plugins) |
|
||||
| AWS (Strands) | **Approach 3** Returns an [AgentResult](https://strandsagents.com/docs/api/python/strands.agent.agent_result/) (Primary) with messages and a reason for the run's completion. | **Approach 2** [Streams events](https://strandsagents.com/docs/api/python/strands.agent.agent/) (Primary+Secondary) including, response text, current_tool_use, even data from "callbacks" (strands plugins) |
|
||||
| LangGraph | **Approach 2** A mixed list of all [messages](https://langchain-ai.github.io/langgraph/agents/run_agents/#output-format) | **Approach 2** A mixed list of all [messages](https://langchain-ai.github.io/langgraph/agents/run_agents/#output-format) |
|
||||
| Agno | **Combination of various approaches** Returns a [RunResponse](https://docs.agno.com/reference/agents/run-response) object with text content, messages (essentially chat history including inputs and instructions), reasoning and thinking text properties. Secondary events could potentially be extracted from messages. | **Approach 2** Returns [RunResponseEvent](https://docs.agno.com/reference/agents/run-response#runresponseevent-types-and-attributes) objects including tool call, memory update, etc, information, where the [RunResponseCompletedEvent](https://docs.agno.com/reference/agents/run-response#runresponsecompletedevent) has similar properties to RunResponse|
|
||||
| A2A | **Approach 3** Returns a [Task or Message](https://a2aproject.github.io/A2A/latest/specification/#71-messagesend) where the message is the final result (Primary) and task is a reference to a long running process. | **Approach 2** Returns a [stream](https://a2aproject.github.io/A2A/latest/specification/#72-messagestream) that contains task updates (Secondary) and a final message (Primary) |
|
||||
@@ -496,7 +496,7 @@ We need to decide what AIContent types, each agent response type will be mapped
|
||||
|-|-|
|
||||
| AutoGen | **Approach 1** Supports [configuring an agent](https://microsoft.github.io/autogen/stable/user-guide/agentchat-user-guide/tutorial/agents.html#structured-output) at agent creation. |
|
||||
| Google ADK | **Approach 1** Both [input and output schemas can be specified for LLM Agents](https://google.github.io/adk-docs/agents/llm-agents/#structuring-data-input_schema-output_schema-output_key) at construction time. This option is specific to this agent type and other agent types do not necessarily support |
|
||||
| AWS (Strands) | **Approach 2** Supports a special invocation method called [structured_output](https://strandsagents.com/docs/user-guide/concepts/agents/structured-output/) |
|
||||
| AWS (Strands) | **Approach 2** Supports a special invocation method called [structured_output](https://strandsagents.com/docs/api/python/strands.agent.agent/) |
|
||||
| LangGraph | **Approach 1** Supports [configuring an agent](https://langchain-ai.github.io/langgraph/agents/agents/?h=structured#6-configure-structured-output) at agent construction time, and a [structured response](https://langchain-ai.github.io/langgraph/agents/run_agents/#output-format) can be retrieved as a special property on the agent response |
|
||||
| Agno | **Approach 1** Supports [configuring an agent](https://docs.agno.com/input-output/structured-output/agent) at agent construction time |
|
||||
| A2A | **Informal Approach 2** Doesn't formally support schema negotiation, but [hints can be provided via metadata](https://a2a-protocol.org/latest/specification/#97-structured-data-exchange-requesting-and-providing-json) at invocation time |
|
||||
@@ -508,7 +508,7 @@ We need to decide what AIContent types, each agent response type will be mapped
|
||||
|-|-|
|
||||
| AutoGen | Supports a [stop reason](https://microsoft.github.io/autogen/stable/reference/python/autogen_agentchat.base.html#autogen_agentchat.base.TaskResult.stop_reason) which is a freeform text string |
|
||||
| Google ADK | [No equivalent present](https://github.com/google/adk-python/blob/main/src/google/adk/events/event.py) |
|
||||
| AWS (Strands) | Exposes a `stop_reason` property on the [AgentResult](https://strandsagents.com/docs/api/python/strands.agent.agent_result/#agentresult) class with options that are tied closely to LLM operations. |
|
||||
| AWS (Strands) | Exposes a [stop_reason](https://strandsagents.com/docs/api/python/strands.types.event_loop/) property on the [AgentResult](https://strandsagents.com/docs/api/python/strands.agent.agent_result/) class with options that are tied closely to LLM operations. |
|
||||
| LangGraph | No equivalent present, output contains only [messages](https://langchain-ai.github.io/langgraph/agents/run_agents/#output-format) |
|
||||
| Agno | [No equivalent present](https://docs.agno.com/reference/agents/run-response) |
|
||||
| A2A | No equivalent present, response only contains a [message](https://a2a-protocol.org/latest/specification/#64-message-object) or [task](https://a2a-protocol.org/latest/specification/#61-task-object). |
|
||||
|
||||
+14
-5
@@ -17,7 +17,7 @@ internal sealed class Tools(ILogger<Tools> logger)
|
||||
[Description("Starts a content generation workflow and returns the instance ID for tracking.")]
|
||||
public string StartContentGenerationWorkflow([Description("The topic for content generation")] string topic)
|
||||
{
|
||||
this._logger.LogInformation("Starting content generation workflow for topic: {Topic}", topic);
|
||||
this._logger.LogInformation("Starting content generation workflow for topic: {Topic}", SanitizeLogValue(topic));
|
||||
|
||||
const int MaxReviewAttempts = 3;
|
||||
const float ApprovalTimeoutHours = 72;
|
||||
@@ -34,7 +34,7 @@ internal sealed class Tools(ILogger<Tools> logger)
|
||||
|
||||
this._logger.LogInformation(
|
||||
"Content generation workflow scheduled to be started for topic '{Topic}' with instance ID: {InstanceId}",
|
||||
topic,
|
||||
SanitizeLogValue(topic),
|
||||
instanceId);
|
||||
|
||||
return $"Workflow started with instance ID: {instanceId}";
|
||||
@@ -45,7 +45,7 @@ internal sealed class Tools(ILogger<Tools> logger)
|
||||
[Description("The instance ID of the workflow to check")] string instanceId,
|
||||
[Description("Whether to include detailed information")] bool includeDetails = true)
|
||||
{
|
||||
this._logger.LogInformation("Getting status for workflow instance: {InstanceId}", instanceId);
|
||||
this._logger.LogInformation("Getting status for workflow instance: {InstanceId}", SanitizeLogValue(instanceId));
|
||||
|
||||
// Get the current agent context using the session-static property
|
||||
OrchestrationMetadata? status = await DurableAgentContext.Current.GetOrchestrationStatusAsync(
|
||||
@@ -54,7 +54,7 @@ internal sealed class Tools(ILogger<Tools> logger)
|
||||
|
||||
if (status is null)
|
||||
{
|
||||
this._logger.LogInformation("Workflow instance '{InstanceId}' not found.", instanceId);
|
||||
this._logger.LogInformation("Workflow instance '{InstanceId}' not found.", SanitizeLogValue(instanceId));
|
||||
return new
|
||||
{
|
||||
instanceId,
|
||||
@@ -78,7 +78,16 @@ internal sealed class Tools(ILogger<Tools> logger)
|
||||
[Description("The instance ID of the workflow to submit feedback for")] string instanceId,
|
||||
[Description("Feedback to submit")] HumanApprovalResponse feedback)
|
||||
{
|
||||
this._logger.LogInformation("Submitting human approval for workflow instance: {InstanceId}", instanceId);
|
||||
this._logger.LogInformation("Submitting human approval for workflow instance: {InstanceId}", SanitizeLogValue(instanceId));
|
||||
await DurableAgentContext.Current.RaiseOrchestrationEventAsync(instanceId, "HumanApproval", feedback);
|
||||
}
|
||||
|
||||
/// <summary>
|
||||
/// Sanitizes a user-provided value for safe inclusion in log entries
|
||||
/// by removing control characters that could be used for log forging.
|
||||
/// </summary>
|
||||
private static string SanitizeLogValue(string value) =>
|
||||
value
|
||||
.Replace("\r", string.Empty, StringComparison.Ordinal)
|
||||
.Replace("\n", string.Empty, StringComparison.Ordinal);
|
||||
}
|
||||
|
||||
+20
-4
@@ -157,8 +157,8 @@ public sealed class FunctionTriggers
|
||||
|
||||
this._logger.LogInformation(
|
||||
"Resuming stream for conversation {ConversationId} from cursor: {Cursor}",
|
||||
conversationId,
|
||||
cursor ?? "(beginning)");
|
||||
SanitizeLogValue(conversationId),
|
||||
SanitizeLogValue(cursor) ?? "(beginning)");
|
||||
|
||||
// Check Accept header to determine response format
|
||||
// text/plain = raw text output (ideal for terminals)
|
||||
@@ -205,7 +205,7 @@ public sealed class FunctionTriggers
|
||||
{
|
||||
if (chunk.Error != null)
|
||||
{
|
||||
this._logger.LogWarning("Stream error for conversation {ConversationId}: {Error}", conversationId, chunk.Error);
|
||||
this._logger.LogWarning("Stream error for conversation {ConversationId}: {Error}", SanitizeLogValue(conversationId), chunk.Error);
|
||||
await WriteErrorAsync(httpContext.Response, chunk.Error, useSseFormat, cancellationToken);
|
||||
break;
|
||||
}
|
||||
@@ -224,7 +224,7 @@ public sealed class FunctionTriggers
|
||||
}
|
||||
catch (OperationCanceledException)
|
||||
{
|
||||
this._logger.LogInformation("Client disconnected from stream {ConversationId}", conversationId);
|
||||
this._logger.LogInformation("Client disconnected from stream {ConversationId}", SanitizeLogValue(conversationId));
|
||||
}
|
||||
|
||||
return new EmptyResult();
|
||||
@@ -316,4 +316,20 @@ public sealed class FunctionTriggers
|
||||
|
||||
await response.WriteAsync(sb.ToString());
|
||||
}
|
||||
|
||||
/// <summary>
|
||||
/// Sanitizes a user-provided value for safe inclusion in log entries
|
||||
/// by removing control characters that could be used for log forging.
|
||||
/// </summary>
|
||||
private static string? SanitizeLogValue(string? value)
|
||||
{
|
||||
if (value is null)
|
||||
{
|
||||
return null;
|
||||
}
|
||||
|
||||
return value
|
||||
.Replace("\r", string.Empty, StringComparison.Ordinal)
|
||||
.Replace("\n", string.Empty, StringComparison.Ordinal);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -4,6 +4,9 @@
|
||||
// In this case the OpenAI responses service will invoke any MCP tools as required. MCP tools are not invoked by the Agent Framework.
|
||||
// The sample demonstrates how to use MCP tools with auto approval by setting ApprovalMode to NeverRequire.
|
||||
|
||||
#pragma warning disable MEAI001 // HostedMcpServerTool, HostedMcpServerToolApprovalMode are experimental
|
||||
#pragma warning disable OPENAI001 // GetResponsesClient is experimental
|
||||
|
||||
using Azure.AI.AgentServer.AgentFramework.Extensions;
|
||||
using Azure.AI.OpenAI;
|
||||
using Azure.Identity;
|
||||
|
||||
@@ -3,6 +3,8 @@
|
||||
// This sample demonstrates a multi-agent workflow with Writer and Reviewer agents
|
||||
// using Azure AI Foundry AIProjectClient and the Agent Framework WorkflowBuilder.
|
||||
|
||||
#pragma warning disable CA2252 // AIProjectClient and Agents API require opting into preview features
|
||||
|
||||
using Azure.AI.AgentServer.AgentFramework.Extensions;
|
||||
using Azure.AI.Projects;
|
||||
using Azure.Identity;
|
||||
|
||||
@@ -4,6 +4,8 @@
|
||||
// Uses Microsoft Agent Framework with Azure AI Foundry.
|
||||
// Ready for deployment to Foundry Hosted Agent service.
|
||||
|
||||
#pragma warning disable CA2252 // AIProjectClient and Agents API require opting into preview features
|
||||
|
||||
using System.ComponentModel;
|
||||
using System.Globalization;
|
||||
using System.Text;
|
||||
|
||||
@@ -4,6 +4,12 @@
|
||||
|
||||
### Changed
|
||||
|
||||
- Filter empty `AIContent` from durable agent state responses ([#4670](https://github.com/microsoft/agent-framework/pull/4670))
|
||||
|
||||
## v1.0.0-preview.260311.1
|
||||
|
||||
### Changed
|
||||
|
||||
- Added TTL configuration for durable agent entities ([#2679](https://github.com/microsoft/agent-framework/pull/2679))
|
||||
- Switch to new "Run" method name ([#2843](https://github.com/microsoft/agent-framework/pull/2843))
|
||||
- Removed AgentThreadMetadata and used AgentSessionId directly instead ([#3067](https://github.com/microsoft/agent-framework/pull/3067));
|
||||
@@ -16,6 +22,8 @@
|
||||
- Marked all `RunAsync<T>` overloads as `new`, added missing ones, and added support for primitives and arrays ([#3803](https://github.com/microsoft/agent-framework/pull/3803))
|
||||
- Improve session cast error message quality and consistency ([#3973](https://github.com/microsoft/agent-framework/pull/3973))
|
||||
|
||||
NOTE: Some of the above changes may have been part of earlier releases not mentioned in this file.
|
||||
|
||||
## v1.0.0-preview.251204.1
|
||||
|
||||
- Added orchestration ID to durable agent entity state ([#2137](https://github.com/microsoft/agent-framework/pull/2137))
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
// Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
using System.Text.Json.Serialization;
|
||||
using Microsoft.Extensions.AI;
|
||||
|
||||
namespace Microsoft.Agents.AI.DurableTask.State;
|
||||
|
||||
@@ -28,7 +29,10 @@ internal sealed class DurableAgentStateResponse : DurableAgentStateEntry
|
||||
{
|
||||
CorrelationId = correlationId,
|
||||
CreatedAt = response.CreatedAt ?? response.Messages.Max(m => m.CreatedAt) ?? DateTimeOffset.UtcNow,
|
||||
Messages = response.Messages.Select(DurableAgentStateMessage.FromChatMessage).ToList(),
|
||||
Messages = response.Messages
|
||||
.Where(HasSerializableContent)
|
||||
.Select(DurableAgentStateMessage.FromChatMessage)
|
||||
.ToList(),
|
||||
Usage = DurableAgentStateUsage.FromUsage(response.Usage)
|
||||
};
|
||||
}
|
||||
@@ -46,4 +50,18 @@ internal sealed class DurableAgentStateResponse : DurableAgentStateEntry
|
||||
Usage = this.Usage?.ToUsageDetails(),
|
||||
};
|
||||
}
|
||||
|
||||
// Checks whether a ChatMessage has any content that will produce meaningful serialized data.
|
||||
// Known derived AIContent types (TextContent, FunctionCallContent, etc.) are always serializable.
|
||||
// Base AIContent instances only carry RawRepresentation (which is [JsonIgnore]), Annotations, and
|
||||
// AdditionalProperties. We keep the message if any base AIContent has annotations or additional
|
||||
// properties set. NOTE: if AIContent gains new serializable properties in the future, this check
|
||||
// should be updated accordingly.
|
||||
private static bool HasSerializableContent(ChatMessage message)
|
||||
{
|
||||
return message.Contents.Any(c =>
|
||||
c.GetType() != typeof(AIContent) ||
|
||||
c.Annotations?.Count > 0 ||
|
||||
c.AdditionalProperties?.Count > 0);
|
||||
}
|
||||
}
|
||||
|
||||
+142
@@ -0,0 +1,142 @@
|
||||
// Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
using Microsoft.Agents.AI.DurableTask.State;
|
||||
using Microsoft.Extensions.AI;
|
||||
|
||||
namespace Microsoft.Agents.AI.DurableTask.Tests.Unit.State;
|
||||
|
||||
public sealed class DurableAgentStateResponseTests
|
||||
{
|
||||
[Fact]
|
||||
public void FromResponseDropsMessagesContainingOnlyOpaqueContent()
|
||||
{
|
||||
// Arrange: one message with real text, one with only opaque AIContent
|
||||
ChatMessage usefulMessage = new(ChatRole.Assistant, "Hello, world!")
|
||||
{
|
||||
CreatedAt = DateTimeOffset.UtcNow
|
||||
};
|
||||
ChatMessage opaqueOnlyMessage = new(ChatRole.Assistant, [
|
||||
new AIContent
|
||||
{
|
||||
RawRepresentation = new { kind = "sessionEvent", sessionId = "s123" }
|
||||
}])
|
||||
{
|
||||
CreatedAt = DateTimeOffset.UtcNow.AddSeconds(1)
|
||||
};
|
||||
|
||||
AgentResponse response = new(new List<ChatMessage> { usefulMessage, opaqueOnlyMessage })
|
||||
{
|
||||
CreatedAt = DateTimeOffset.UtcNow
|
||||
};
|
||||
|
||||
// Act
|
||||
DurableAgentStateResponse durableResponse = DurableAgentStateResponse.FromResponse("corr-123", response);
|
||||
|
||||
// Assert: only the useful message survives
|
||||
DurableAgentStateMessage durableMessage = Assert.Single(durableResponse.Messages);
|
||||
Assert.Equal(ChatRole.Assistant.Value, durableMessage.Role);
|
||||
|
||||
// Round-trip to verify the content is correct
|
||||
AgentResponse convertedResponse = durableResponse.ToResponse();
|
||||
ChatMessage convertedMessage = Assert.Single(convertedResponse.Messages);
|
||||
TextContent textContent = Assert.IsType<TextContent>(Assert.Single(convertedMessage.Contents));
|
||||
Assert.Equal("Hello, world!", textContent.Text);
|
||||
}
|
||||
|
||||
[Fact]
|
||||
public void FromResponseKeepsMessagesWithMixedContent()
|
||||
{
|
||||
// Arrange: one message with both real text and opaque AIContent
|
||||
ChatMessage mixedMessage = new(ChatRole.Assistant, [
|
||||
new TextContent("Some useful text"),
|
||||
new AIContent { RawRepresentation = new { kind = "metadata" } }])
|
||||
{
|
||||
CreatedAt = DateTimeOffset.UtcNow
|
||||
};
|
||||
|
||||
AgentResponse response = new(new List<ChatMessage> { mixedMessage })
|
||||
{
|
||||
CreatedAt = DateTimeOffset.UtcNow
|
||||
};
|
||||
|
||||
// Act
|
||||
DurableAgentStateResponse durableResponse = DurableAgentStateResponse.FromResponse("corr-456", response);
|
||||
|
||||
// Assert: the message is kept because it contains at least one serializable content
|
||||
DurableAgentStateMessage durableMessage = Assert.Single(durableResponse.Messages);
|
||||
Assert.Equal(ChatRole.Assistant.Value, durableMessage.Role);
|
||||
}
|
||||
|
||||
[Fact]
|
||||
public void FromResponseDropsAllMessagesWhenAllAreOpaque()
|
||||
{
|
||||
// Arrange: all messages contain only opaque AIContent
|
||||
ChatMessage opaque1 = new(ChatRole.Assistant, [
|
||||
new AIContent { RawRepresentation = new { kind = "event1" } }])
|
||||
{
|
||||
CreatedAt = DateTimeOffset.UtcNow
|
||||
};
|
||||
ChatMessage opaque2 = new(ChatRole.Assistant, [
|
||||
new AIContent { RawRepresentation = new { kind = "event2" } }])
|
||||
{
|
||||
CreatedAt = DateTimeOffset.UtcNow.AddSeconds(1)
|
||||
};
|
||||
|
||||
AgentResponse response = new(new List<ChatMessage> { opaque1, opaque2 })
|
||||
{
|
||||
CreatedAt = DateTimeOffset.UtcNow
|
||||
};
|
||||
|
||||
// Act
|
||||
DurableAgentStateResponse durableResponse = DurableAgentStateResponse.FromResponse("corr-789", response);
|
||||
|
||||
// Assert: no messages stored
|
||||
Assert.Empty(durableResponse.Messages);
|
||||
}
|
||||
|
||||
[Fact]
|
||||
public void FromResponseKeepsBaseAIContentWithAnnotations()
|
||||
{
|
||||
// Arrange: base AIContent with annotations should be kept
|
||||
AIContent contentWithAnnotations = new()
|
||||
{
|
||||
RawRepresentation = new { kind = "event" },
|
||||
Annotations = [new AIAnnotation() { AdditionalProperties = new() { ["cite"] = "ref-1" } }]
|
||||
};
|
||||
ChatMessage message = new(ChatRole.Assistant, [contentWithAnnotations])
|
||||
{
|
||||
CreatedAt = DateTimeOffset.UtcNow
|
||||
};
|
||||
|
||||
AgentResponse response = new([message]) { CreatedAt = DateTimeOffset.UtcNow };
|
||||
|
||||
// Act
|
||||
DurableAgentStateResponse durableResponse = DurableAgentStateResponse.FromResponse("corr-ann", response);
|
||||
|
||||
// Assert: message is kept because the AIContent has annotations
|
||||
Assert.Single(durableResponse.Messages);
|
||||
}
|
||||
|
||||
[Fact]
|
||||
public void FromResponseKeepsBaseAIContentWithAdditionalProperties()
|
||||
{
|
||||
// Arrange: base AIContent with additional properties should be kept
|
||||
AIContent contentWithProps = new()
|
||||
{
|
||||
RawRepresentation = new { kind = "event" },
|
||||
AdditionalProperties = new() { ["custom_key"] = "custom_value" }
|
||||
};
|
||||
ChatMessage message = new(ChatRole.Assistant, [contentWithProps])
|
||||
{
|
||||
CreatedAt = DateTimeOffset.UtcNow
|
||||
};
|
||||
|
||||
AgentResponse response = new([message]) { CreatedAt = DateTimeOffset.UtcNow };
|
||||
|
||||
// Act
|
||||
DurableAgentStateResponse durableResponse = DurableAgentStateResponse.FromResponse("corr-props", response);
|
||||
|
||||
// Assert: message is kept because the AIContent has additional properties
|
||||
Assert.Single(durableResponse.Messages);
|
||||
}
|
||||
}
|
||||
+4
-4
@@ -69,7 +69,7 @@ def equal(arg1: str, arg2: str) -> bool:
|
||||
|
||||
```python
|
||||
# Core
|
||||
from agent_framework import ChatAgent, Message, tool
|
||||
from agent_framework import Agent, Message, tool
|
||||
|
||||
# Components
|
||||
from agent_framework.observability import enable_instrumentation
|
||||
@@ -82,16 +82,16 @@ from agent_framework.azure import AzureOpenAIChatClient
|
||||
## Public API and Exports
|
||||
|
||||
In `__init__.py` files that define package-level public APIs, use direct re-export imports plus an explicit
|
||||
`__all__`. Avoid identity aliases like `from ._agents import ChatAgent as ChatAgent`, and avoid
|
||||
`__all__`. Avoid identity aliases like `from ._agents import Agent as Agent`, and avoid
|
||||
`from module import *`.
|
||||
|
||||
Do not define `__all__` in internal non-`__init__.py` modules. Exception: modules intentionally exposed as a
|
||||
public import surface (for example, `agent_framework.observability`) should define `__all__`.
|
||||
|
||||
```python
|
||||
__all__ = ["ChatAgent", "Message", "ChatResponse"]
|
||||
__all__ = ["Agent", "Message", "ChatResponse"]
|
||||
|
||||
from ._agents import ChatAgent
|
||||
from ._agents import Agent
|
||||
from ._types import Message, ChatResponse
|
||||
```
|
||||
|
||||
|
||||
+43
-1
@@ -33,13 +33,44 @@ Uses [uv](https://github.com/astral-sh/uv) for dependency management and
|
||||
# Full setup (venv + install + prek hooks)
|
||||
uv run poe setup
|
||||
|
||||
# Install/update all dependencies
|
||||
# Install dependencies from lockfile (frozen resolution with prerelease policy)
|
||||
uv run poe install
|
||||
|
||||
# Create venv with specific Python version
|
||||
uv run poe venv --python 3.12
|
||||
|
||||
# Intentionally upgrade a specific dependency to reduce lockfile conflicts
|
||||
uv lock --upgrade-package <dependency-name> && uv run poe install
|
||||
|
||||
# Refresh all dev dependency pins, lockfile, and validation in one run
|
||||
uv run poe upgrade-dev-dependencies
|
||||
|
||||
# First, run workspace-wide lower/upper compatibility gates
|
||||
uv run poe validate-dependency-bounds-test
|
||||
# Defaults to --project "*"; pass a package to scope test mode
|
||||
uv run poe validate-dependency-bounds-test --project <workspace-package-name>
|
||||
|
||||
# Then expand bounds for one dependency in the target package
|
||||
uv run poe validate-dependency-bounds-project --mode both --project <workspace-package-name> --dependency "<dependency-name>"
|
||||
|
||||
# Repo-wide automation can reuse the same task
|
||||
uv run poe validate-dependency-bounds-project --mode upper --project "*"
|
||||
|
||||
# Add a dependency to one project and run both validators for that project/dependency
|
||||
uv run poe add-dependency-and-validate-bounds --project <workspace-package-name> --dependency "<dependency-spec>"
|
||||
```
|
||||
|
||||
### Dependency Bound Notes
|
||||
|
||||
- Stable dependencies (`>=1.0`) should typically be bounded as `>=<known-good>,<next-major>`.
|
||||
- Prerelease (`dev`/`a`/`b`/`rc`) and `<1.0` dependencies should use hard bounds with an explicit upper cap (avoid open-ended ranges).
|
||||
- For `<1.0` dependencies, prefer the broadest validated range the package can really support. That may be a patch line, a minor line, or multiple minor lines when checks/tests show the broader lane is compatible.
|
||||
- Prefer supporting multiple majors when practical; if APIs diverge across supported majors, use version-conditional imports/paths.
|
||||
- For dependency changes, run workspace-wide bound gates first, then `validate-dependency-bounds-project --mode both` for the target package/dependency to keep minimum and maximum constraints current. The same task can also drive repo-wide upper-bound automation by using `--project "*"` and omitting `--dependency`.
|
||||
- Prefer targeted lock updates with `uv lock --upgrade-package <dependency-name>` to reduce `uv.lock` merge conflicts.
|
||||
- Use `add-dependency-and-validate-bounds` for package-scoped dependency additions plus bound validation in one command.
|
||||
- Use `upgrade-dev-dependencies` for repo-wide dev tooling refreshes; it repins dev dependencies, refreshes `uv.lock`, and reruns `check`, `typing`, and `test`.
|
||||
|
||||
## Lazy Loading Pattern
|
||||
|
||||
Provider folders in core use `__getattr__` to lazy load from connector packages:
|
||||
@@ -74,6 +105,17 @@ def __getattr__(name: str) -> Any:
|
||||
4. Do **NOT** add to `[all]` extra in `packages/core/pyproject.toml`
|
||||
5. Do **NOT** create lazy loading in core yet
|
||||
|
||||
Recommended dependency workflow during connector implementation:
|
||||
|
||||
1. Add the dependency to the target package:
|
||||
`uv run poe add-dependency-to-project --project <workspace-package-name> --dependency "<dependency-spec>"`
|
||||
2. Implement connector code and tests.
|
||||
3. Validate dependency bounds for that package/dependency:
|
||||
`uv run poe validate-dependency-bounds-project --mode both --project <workspace-package-name> --dependency "<dependency-name>"`
|
||||
4. If the package has meaningful tests/checks that validate dependency compatibility, you can use the add + validation flow in one command:
|
||||
`uv run poe add-dependency-and-validate-bounds --project <workspace-package-name> --dependency "<dependency-spec>"`
|
||||
If compatibility checks are not in place yet, add the dependency first, then implement tests before running bound validation.
|
||||
|
||||
### Promotion to Stable
|
||||
|
||||
1. Move samples to root `samples/` folder
|
||||
|
||||
@@ -127,7 +127,12 @@ def create_agent(name: str, tool_mode: Literal['auto', 'required', 'none'] | Cha
|
||||
Avoid `**kwargs` unless absolutely necessary. It should only be used as an escape route, not for well-known flows of data:
|
||||
|
||||
- **Prefer named parameters**: If there are known extra arguments being passed, use explicit named parameters instead of kwargs
|
||||
- **Prefer purpose-specific buckets over generic kwargs**: If a flexible payload is still needed, use an explicit named parameter such as `additional_properties`, `function_invocation_kwargs`, or `client_kwargs` rather than a blanket `**kwargs`
|
||||
- **Subclassing support**: kwargs is acceptable in methods that are part of classes designed for subclassing, allowing subclass-defined kwargs to pass through without issues. In this case, clearly document that kwargs exists for subclass extensibility and not for passing arbitrary data
|
||||
- **Make known flows explicit first**: For abstract hooks, move known data flows into explicit parameters before leaving `**kwargs` behind for subclass extensibility (for example, prefer `state=` explicitly instead of passing it through kwargs)
|
||||
- **Prefer explicit metadata containers**: For constructors that expose metadata, prefer an explicit `additional_properties` parameter.
|
||||
- **Keep SDK passthroughs narrow and documented**: A kwargs escape hatch may be acceptable for provider helper APIs that pass through to a large or unstable external SDK surface, but it should be documented as SDK passthrough and revisited regularly
|
||||
- **Do not keep passthrough kwargs on wrappers that do not use them**: Convenience wrappers and session helpers should not accept generic kwargs merely to forward or ignore them
|
||||
- **Remove when possible**: In other cases, removing kwargs is likely better than keeping it
|
||||
- **Separate kwargs by purpose**: When combining kwargs for multiple purposes, use specific parameters like `client_kwargs: dict[str, Any]` instead of mixing everything in `**kwargs`
|
||||
- **Always document**: If kwargs must be used, always document how it's used, either by referencing external documentation or explaining its purpose
|
||||
@@ -160,10 +165,14 @@ user_msg = Message("user", ["Hello, world!"])
|
||||
asst_msg = Message("assistant", ["Hello, world!"])
|
||||
|
||||
# ❌ Not preferred - unnecessary inheritance
|
||||
from agent_framework import UserMessage, AssistantMessage
|
||||
class UserMessage(Message):
|
||||
pass
|
||||
|
||||
user_msg = UserMessage(content="Hello, world!")
|
||||
asst_msg = AssistantMessage(content="Hello, world!")
|
||||
class AssistantMessage(Message):
|
||||
pass
|
||||
|
||||
user_msg = UserMessage("user", ["Hello, world!"])
|
||||
asst_msg = AssistantMessage("assistant", ["Hello, world!"])
|
||||
```
|
||||
|
||||
### Import Structure
|
||||
@@ -383,6 +392,19 @@ All non-core packages declare a lower bound on `agent-framework-core` (e.g., `"a
|
||||
- **Core version changes**: When `agent-framework-core` is updated with breaking or significant changes and its version is bumped, update the `agent-framework-core>=...` lower bound in every other package's `pyproject.toml` to match the new core version.
|
||||
- **Non-core version changes**: Non-core packages (connectors, extensions) can have their own versions incremented independently while keeping the existing core lower bound pinned. Only raise the core lower bound if the non-core package actually depends on new core APIs.
|
||||
|
||||
### External Dependency Version Bounds
|
||||
|
||||
The guiding principle for external dependencies is to make the range of allowed versions as broad as possible, even if that means we have to do some conditional imports, and other tricks to allow small changes in versions.
|
||||
So we use bounded ranges for external package dependencies in `pyproject.toml`:
|
||||
|
||||
|
||||
- For stable dependencies (`>=1.0.0`), use a lower bound at a known-good version and an explicit upper bound that reflects the maximum major version we currently support (for example: `openai>=1.99.0,<3`).
|
||||
- For prerelease (`dev`/`a`/`b`/`rc`) dependencies, use a known-good lower bound with a hard upper boundary in the same prerelease line (for example: `azure-ai-projects>=2.0.0b3,<2.0.0b4`).
|
||||
- For `<1.0.0` dependencies, use a known-good bounded range with an explicit upper cap. Prefer the broadest validated range the package can actually support: that may be a patch line, a minor line, or multiple minor lines (for example: `a2a-sdk>=0.3.5,<0.4.0`, `fastapi>=0.115.0,<0.136.0`, `uvicorn>=0.30.0,<0.39.0`).
|
||||
- For prerelease (`dev`/`a`/`b`/`rc`) dependencies, use a known-good bounded range with a hard upper cap and keep the range only as broad as the package's validation coverage justifies.
|
||||
- Prefer keeping support for multiple major versions when practical. This may mean that the upper bound spans multiple major versions when the dependency maintains backward compatibility; if APIs differ between supported majors, version-conditional imports/branches are acceptable to preserve compatibility.
|
||||
- When adding or changing an external dependency, first run `uv run poe validate-dependency-bounds-test` to validate workspace-wide lower/upper compatibility, then run `uv run poe validate-dependency-bounds-project --mode both --project <workspace-package-name> --dependency "<dependency-name>"` to expand package-scoped bounds.
|
||||
|
||||
### Installation Options
|
||||
|
||||
Connectors are distributed as separate packages and are not imported by default in the core package. Users install the specific connectors they need:
|
||||
|
||||
+33
-1
@@ -217,10 +217,13 @@ uv run poe setup --python 3.12
|
||||
```
|
||||
|
||||
#### `install`
|
||||
Install all dependencies including extras and dev dependencies, including updates:
|
||||
Install all dependencies (including extras and dev dependencies) from the lockfile using frozen resolution:
|
||||
```bash
|
||||
uv run poe install
|
||||
```
|
||||
For intentional dependency upgrades, run `uv lock --upgrade-package <dependency-name>` and then run `uv run poe install`.
|
||||
|
||||
For repo-wide dev tooling refreshes, run `uv run poe upgrade-dev-dependencies` to repin dev dependencies, refresh `uv.lock`, and rerun validation, typing, and tests.
|
||||
|
||||
#### `venv`
|
||||
Create a virtual environment with specified Python version or switch python version:
|
||||
@@ -278,6 +281,35 @@ Lint markdown code blocks:
|
||||
uv run poe markdown-code-lint
|
||||
```
|
||||
|
||||
#### `validate-dependency-bounds-test`
|
||||
Run workspace-wide dependency compatibility gates at lower and upper resolutions. This runs test + pyright across all packages and stops on first failure:
|
||||
```bash
|
||||
uv run poe validate-dependency-bounds-test
|
||||
# Defaults to --project "*"; pass a package to scope test mode
|
||||
uv run poe validate-dependency-bounds-test --project <workspace-package-name>
|
||||
```
|
||||
|
||||
#### `validate-dependency-bounds-project`
|
||||
Validate and extend dependency bounds for a single dependency in a single package. Use `--mode lower`, `--mode upper`, or the default `--mode both`:
|
||||
```bash
|
||||
uv run poe validate-dependency-bounds-project --mode both --project <workspace-package-name> --dependency "<dependency-name>"
|
||||
```
|
||||
`--project` defaults to `*`, and `--dependency` is optional. Automation can use `--mode upper --project "*"` to run the upper-bound pass across the workspace.
|
||||
For `<1.0` dependencies, prefer the broadest validated range the package can really support. That may still be a single patch or minor line, but multi-minor ranges are fine when the package's checks/tests prove they work.
|
||||
|
||||
#### `add-dependency-and-validate-bounds`
|
||||
Add an external dependency to a workspace project and run both validators for that same project/dependency:
|
||||
```bash
|
||||
uv run poe add-dependency-and-validate-bounds --project <workspace-package-name> --dependency "<dependency-spec>"
|
||||
```
|
||||
|
||||
#### `upgrade-dev-dependencies`
|
||||
Refresh exact dev dependency pins across the workspace, run `uv lock --upgrade`, reinstall from the frozen lockfile, then rerun validation, typing, and tests:
|
||||
```bash
|
||||
uv run poe upgrade-dev-dependencies
|
||||
```
|
||||
Use this for repo-wide dev tooling refreshes. For targeted runtime dependency upgrades, prefer `uv lock --upgrade-package <dependency-name>` plus the package-scoped bound validation tasks above.
|
||||
|
||||
### Comprehensive Checks
|
||||
|
||||
#### `check-packages`
|
||||
|
||||
@@ -6,7 +6,7 @@ import base64
|
||||
import json
|
||||
import re
|
||||
import uuid
|
||||
from collections.abc import AsyncIterable, Awaitable, Sequence
|
||||
from collections.abc import AsyncIterable, Awaitable, Mapping, Sequence
|
||||
from typing import Any, Final, Literal, TypeAlias, overload
|
||||
|
||||
import httpx
|
||||
@@ -226,6 +226,8 @@ class A2AAgent(AgentTelemetryLayer, BaseAgent):
|
||||
*,
|
||||
stream: Literal[False] = ...,
|
||||
session: AgentSession | None = None,
|
||||
function_invocation_kwargs: Mapping[str, Any] | None = None,
|
||||
client_kwargs: Mapping[str, Any] | None = None,
|
||||
continuation_token: A2AContinuationToken | None = None,
|
||||
background: bool = False,
|
||||
**kwargs: Any,
|
||||
@@ -238,17 +240,21 @@ class A2AAgent(AgentTelemetryLayer, BaseAgent):
|
||||
*,
|
||||
stream: Literal[True],
|
||||
session: AgentSession | None = None,
|
||||
function_invocation_kwargs: Mapping[str, Any] | None = None,
|
||||
client_kwargs: Mapping[str, Any] | None = None,
|
||||
continuation_token: A2AContinuationToken | None = None,
|
||||
background: bool = False,
|
||||
**kwargs: Any,
|
||||
) -> ResponseStream[AgentResponseUpdate, AgentResponse[Any]]: ...
|
||||
|
||||
def run(
|
||||
def run( # pyright: ignore[reportIncompatibleMethodOverride]
|
||||
self,
|
||||
messages: AgentRunInputs | None = None,
|
||||
*,
|
||||
stream: bool = False,
|
||||
session: AgentSession | None = None,
|
||||
function_invocation_kwargs: Mapping[str, Any] | None = None,
|
||||
client_kwargs: Mapping[str, Any] | None = None,
|
||||
continuation_token: A2AContinuationToken | None = None,
|
||||
background: bool = False,
|
||||
**kwargs: Any,
|
||||
@@ -261,17 +267,23 @@ class A2AAgent(AgentTelemetryLayer, BaseAgent):
|
||||
Keyword Args:
|
||||
stream: Whether to stream the response. Defaults to False.
|
||||
session: The conversation session associated with the message(s).
|
||||
function_invocation_kwargs: Present for compatibility with the shared agent interface.
|
||||
A2AAgent does not use these values directly.
|
||||
client_kwargs: Present for compatibility with the shared agent interface.
|
||||
A2AAgent does not use these values directly.
|
||||
kwargs: Additional compatibility keyword arguments.
|
||||
A2AAgent does not use these values directly.
|
||||
continuation_token: Optional token to resume a long-running task
|
||||
instead of starting a new one.
|
||||
background: When True, in-progress task updates surface continuation
|
||||
tokens so the caller can poll or resubscribe later. When False
|
||||
(default), the agent internally waits for the task to complete.
|
||||
kwargs: Additional keyword arguments.
|
||||
|
||||
Returns:
|
||||
When stream=False: An Awaitable[AgentResponse].
|
||||
When stream=True: A ResponseStream of AgentResponseUpdate items.
|
||||
"""
|
||||
del function_invocation_kwargs, client_kwargs, kwargs
|
||||
if continuation_token is not None:
|
||||
a2a_stream: AsyncIterable[A2AStreamItem] = self.client.resubscribe(
|
||||
TaskIdParams(id=continuation_token["task_id"])
|
||||
|
||||
@@ -24,7 +24,7 @@ classifiers = [
|
||||
]
|
||||
dependencies = [
|
||||
"agent-framework-core>=1.0.0rc4",
|
||||
"a2a-sdk>=0.3.5",
|
||||
"a2a-sdk>=0.3.5,<0.3.24",
|
||||
]
|
||||
|
||||
[tool.uv]
|
||||
@@ -87,7 +87,7 @@ include = "../../shared_tasks.toml"
|
||||
|
||||
[tool.poe.tasks]
|
||||
mypy = "mypy --config-file $POE_ROOT/pyproject.toml agent_framework_a2a"
|
||||
test = "pytest -m \"not integration\" --cov=agent_framework_a2a --cov-report=term-missing:skip-covered tests"
|
||||
test = 'pytest -m "not integration" --cov=agent_framework_a2a --cov-report=term-missing:skip-covered tests'
|
||||
|
||||
[build-system]
|
||||
requires = ["flit-core >= 3.11,<4.0"]
|
||||
|
||||
@@ -220,7 +220,6 @@ class AGUIChatClient(
|
||||
additional_properties: dict[str, Any] | None = None,
|
||||
middleware: Sequence[ChatAndFunctionMiddlewareTypes] | None = None,
|
||||
function_invocation_configuration: FunctionInvocationConfiguration | None = None,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
"""Initialize the AG-UI chat client.
|
||||
|
||||
@@ -231,13 +230,11 @@ class AGUIChatClient(
|
||||
additional_properties: Additional properties to store
|
||||
middleware: Optional middleware to apply to the client.
|
||||
function_invocation_configuration: Optional function invocation configuration override.
|
||||
**kwargs: Additional arguments passed to BaseChatClient
|
||||
"""
|
||||
super().__init__(
|
||||
additional_properties=additional_properties,
|
||||
middleware=middleware,
|
||||
function_invocation_configuration=function_invocation_configuration,
|
||||
**kwargs,
|
||||
)
|
||||
self._http_service = AGUIHttpService(
|
||||
endpoint=endpoint,
|
||||
|
||||
@@ -8,6 +8,7 @@ import logging
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from agent_framework import BaseChatClient
|
||||
from agent_framework._tools import _append_unique_tools # pyright: ignore[reportPrivateUsage]
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from agent_framework import SupportsAgentRun
|
||||
@@ -22,7 +23,7 @@ def _collect_mcp_tool_functions(mcp_tools: list[Any]) -> list[Any]:
|
||||
mcp_tools: List of MCP tool instances.
|
||||
|
||||
Returns:
|
||||
List of functions from connected MCP tools.
|
||||
Functions from connected MCP tools.
|
||||
"""
|
||||
functions: list[Any] = []
|
||||
for mcp_tool in mcp_tools:
|
||||
@@ -56,7 +57,11 @@ def collect_server_tools(agent: SupportsAgentRun) -> list[Any]:
|
||||
# Include functions from connected MCP tools (only available on Agent)
|
||||
mcp_tools = getattr(agent, "mcp_tools", None)
|
||||
if mcp_tools:
|
||||
server_tools.extend(_collect_mcp_tool_functions(mcp_tools))
|
||||
_append_unique_tools(
|
||||
server_tools,
|
||||
_collect_mcp_tool_functions(mcp_tools),
|
||||
duplicate_error_message="Tool names must be unique. Consider setting `tool_name_prefix` on the MCPTool.",
|
||||
)
|
||||
|
||||
logger.info(f"[TOOLS] Agent has {len(server_tools)} configured tools")
|
||||
for tool in server_tools:
|
||||
@@ -109,26 +114,13 @@ def merge_tools(server_tools: list[Any], client_tools: list[Any] | None) -> list
|
||||
logger.info("[TOOLS] No client tools - not passing tools= parameter (using agent's configured tools)")
|
||||
return None
|
||||
|
||||
server_tool_names = {getattr(tool, "name", None) for tool in server_tools}
|
||||
unique_client_tools = [tool for tool in client_tools if getattr(tool, "name", None) not in server_tool_names]
|
||||
|
||||
if not unique_client_tools:
|
||||
# Same check: must pass server tools if any require approval
|
||||
if server_tools and _has_approval_tools(server_tools):
|
||||
logger.info(
|
||||
f"[TOOLS] Client tools duplicate server but server has approval tools - "
|
||||
f"passing {len(server_tools)} server tools for approval mode"
|
||||
)
|
||||
return server_tools
|
||||
logger.info("[TOOLS] All client tools duplicate server tools - not passing tools= parameter")
|
||||
return None
|
||||
|
||||
combined_tools: list[Any] = []
|
||||
if server_tools:
|
||||
combined_tools.extend(server_tools)
|
||||
combined_tools.extend(unique_client_tools)
|
||||
combined_tools = _append_unique_tools(
|
||||
list(server_tools),
|
||||
client_tools,
|
||||
duplicate_error_message="Tool names must be unique.",
|
||||
)
|
||||
logger.info(
|
||||
f"[TOOLS] Passing tools= parameter with {len(combined_tools)} tools "
|
||||
f"({len(server_tools)} server + {len(unique_client_tools)} unique client)"
|
||||
f"({len(server_tools)} server + {len(client_tools)} client)"
|
||||
)
|
||||
return combined_tools
|
||||
|
||||
@@ -6,13 +6,12 @@ from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import os
|
||||
from typing import cast
|
||||
from typing import Any, cast
|
||||
|
||||
import uvicorn
|
||||
from agent_framework import ChatOptions
|
||||
from agent_framework._clients import SupportsChatGetResponse
|
||||
from agent_framework.ag_ui import add_agent_framework_fastapi_endpoint
|
||||
from agent_framework.anthropic import AnthropicClient
|
||||
from agent_framework.azure import AzureOpenAIChatClient
|
||||
from fastapi import FastAPI
|
||||
from fastapi.middleware.cors import CORSMiddleware
|
||||
@@ -26,6 +25,15 @@ from ..agents.task_steps_agent import task_steps_agent_wrapped
|
||||
from ..agents.ui_generator_agent import ui_generator_agent
|
||||
from ..agents.weather_agent import weather_agent
|
||||
|
||||
AnthropicClient: type[Any] | None
|
||||
try:
|
||||
import agent_framework.anthropic as _anthropic_namespace
|
||||
except ImportError:
|
||||
# If the Anthropic client isn't installed, we can still run the server with Azure OpenAI as the default chat client
|
||||
AnthropicClient = None
|
||||
else:
|
||||
AnthropicClient = cast(type[Any] | None, getattr(_anthropic_namespace, "AnthropicClient", None))
|
||||
|
||||
# Configure logging to file and console (disabled by default - set ENABLE_DEBUG_LOGGING=1 to enable)
|
||||
if os.getenv("ENABLE_DEBUG_LOGGING"):
|
||||
log_file = os.path.join(os.path.dirname(os.path.abspath(__file__)), "..", "..", "ag_ui_server.log")
|
||||
@@ -70,7 +78,9 @@ app.add_middleware(
|
||||
# Set CHAT_CLIENT=anthropic to use Anthropic, defaults to Azure OpenAI
|
||||
client: SupportsChatGetResponse[ChatOptions] = cast(
|
||||
SupportsChatGetResponse[ChatOptions],
|
||||
AnthropicClient() if os.getenv("CHAT_CLIENT", "").lower() == "anthropic" else AzureOpenAIChatClient(),
|
||||
AnthropicClient()
|
||||
if AnthropicClient is not None and os.getenv("CHAT_CLIENT", "").lower() == "anthropic"
|
||||
else AzureOpenAIChatClient(),
|
||||
)
|
||||
|
||||
# Agentic Chat - basic chat agent
|
||||
|
||||
@@ -23,15 +23,15 @@ classifiers = [
|
||||
]
|
||||
dependencies = [
|
||||
"agent-framework-core>=1.0.0rc4",
|
||||
"ag-ui-protocol>=0.1.9",
|
||||
"fastapi>=0.115.0",
|
||||
"uvicorn>=0.30.0"
|
||||
"ag-ui-protocol==0.1.13",
|
||||
"fastapi>=0.115.0,<0.133.1",
|
||||
"uvicorn[standard]>=0.30.0,<0.42.0"
|
||||
]
|
||||
|
||||
[project.optional-dependencies]
|
||||
dev = [
|
||||
"pytest>=8.0.0",
|
||||
"httpx>=0.27.0",
|
||||
"pytest==9.0.2",
|
||||
"httpx==0.28.1",
|
||||
]
|
||||
|
||||
[build-system]
|
||||
@@ -74,4 +74,4 @@ include = "../../shared_tasks.toml"
|
||||
|
||||
[tool.poe.tasks]
|
||||
mypy = "mypy --config-file $POE_ROOT/pyproject.toml agent_framework_ag_ui"
|
||||
test = "pytest -m \"not integration\" --cov=agent_framework_ag_ui --cov-report=term-missing:skip-covered -n auto --dist worksteal tests/ag_ui"
|
||||
test = 'pytest -m "not integration" --cov=agent_framework_ag_ui --cov-report=term-missing:skip-covered -n auto --dist worksteal tests/ag_ui'
|
||||
|
||||
@@ -98,7 +98,11 @@ class StreamingChatClientStub(
|
||||
options: OptionsCoT | ChatOptions[Any] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Awaitable[ChatResponse[Any]] | ResponseStream[ChatResponseUpdate, ChatResponse[Any]]:
|
||||
self.last_session = kwargs.get("session")
|
||||
client_kwargs = kwargs.get("client_kwargs")
|
||||
if isinstance(client_kwargs, Mapping):
|
||||
self.last_session = cast(AgentSession | None, client_kwargs.get("session"))
|
||||
else:
|
||||
self.last_session = None
|
||||
self.last_service_session_id = self.last_session.service_session_id if self.last_session else None
|
||||
return cast(
|
||||
Awaitable[ChatResponse[Any]] | ResponseStream[ChatResponseUpdate, ChatResponse[Any]],
|
||||
|
||||
@@ -702,14 +702,9 @@ async def test_agent_with_use_service_session_is_true(streaming_chat_client_stub
|
||||
"""Test that when use_service_session is True, the AgentSession used to run the agent is set to the service session ID."""
|
||||
from agent_framework.ag_ui import AgentFrameworkAgent
|
||||
|
||||
request_service_session_id: str | None = None
|
||||
|
||||
async def stream_fn(
|
||||
messages: MutableSequence[Message], chat_options: ChatOptions, **kwargs: Any
|
||||
) -> AsyncIterator[ChatResponseUpdate]:
|
||||
nonlocal request_service_session_id
|
||||
session = kwargs.get("session")
|
||||
request_service_session_id = session.service_session_id if session else None
|
||||
yield ChatResponseUpdate(
|
||||
contents=[Content.from_text(text="Response")], response_id="resp_67890", conversation_id="conv_12345"
|
||||
)
|
||||
@@ -719,11 +714,22 @@ async def test_agent_with_use_service_session_is_true(streaming_chat_client_stub
|
||||
|
||||
input_data = {"messages": [{"role": "user", "content": "Hi"}], "thread_id": "conv_123456"}
|
||||
|
||||
# Spy on agent.run to capture the session kwarg at call time (before streaming mutates it)
|
||||
captured_service_session_id: str | None = None
|
||||
original_run = agent.run
|
||||
|
||||
def capturing_run(*args: Any, **kwargs: Any) -> Any:
|
||||
nonlocal captured_service_session_id
|
||||
session = kwargs.get("session")
|
||||
captured_service_session_id = session.service_session_id if session else None
|
||||
return original_run(*args, **kwargs)
|
||||
|
||||
agent.run = capturing_run # type: ignore[assignment, method-assign]
|
||||
|
||||
events: list[Any] = []
|
||||
async for event in wrapper.run(input_data):
|
||||
events.append(event)
|
||||
request_service_session_id = agent.client.last_service_session_id
|
||||
assert request_service_session_id == "conv_123456" # type: ignore[attr-defined] (service_session_id should be set)
|
||||
assert captured_service_session_id == "conv_123456"
|
||||
|
||||
|
||||
async def test_function_approval_mode_executes_tool(streaming_chat_client_stub):
|
||||
|
||||
@@ -2,6 +2,7 @@
|
||||
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
from agent_framework import Agent, tool
|
||||
|
||||
from agent_framework_ag_ui._orchestration._tooling import (
|
||||
@@ -20,7 +21,8 @@ class DummyTool:
|
||||
class MockMCPTool:
|
||||
"""Mock MCP tool that simulates connected MCP tool with functions."""
|
||||
|
||||
def __init__(self, functions: list[DummyTool], is_connected: bool = True) -> None:
|
||||
def __init__(self, functions: list[DummyTool], is_connected: bool = True, name: str = "mock-mcp") -> None:
|
||||
self.name = name
|
||||
self.functions = functions
|
||||
self.is_connected = is_connected
|
||||
|
||||
@@ -45,11 +47,8 @@ def test_merge_tools_filters_duplicates() -> None:
|
||||
server = [DummyTool("a"), DummyTool("b")]
|
||||
client = [DummyTool("b"), DummyTool("c")]
|
||||
|
||||
merged = merge_tools(server, client)
|
||||
|
||||
assert merged is not None
|
||||
names = [getattr(t, "name", None) for t in merged]
|
||||
assert names == ["a", "b", "c"]
|
||||
with pytest.raises(ValueError, match="Duplicate tool name 'b'"):
|
||||
merge_tools(server, client)
|
||||
|
||||
|
||||
def test_register_additional_client_tools_assigns_when_configured() -> None:
|
||||
@@ -131,6 +130,17 @@ def test_collect_server_tools_with_mcp_tools_via_public_property() -> None:
|
||||
assert len(tools) == 2
|
||||
|
||||
|
||||
def test_collect_server_tools_raises_on_duplicate_agent_and_mcp_tool_names() -> None:
|
||||
duplicate_tool = DummyTool("regular_tool")
|
||||
mock_mcp = MockMCPTool([duplicate_tool], is_connected=True, name="docs-mcp")
|
||||
|
||||
agent = _create_chat_agent_with_tool("regular_tool")
|
||||
agent.mcp_tools = [mock_mcp]
|
||||
|
||||
with pytest.raises(ValueError, match="Duplicate tool name 'regular_tool'"):
|
||||
collect_server_tools(agent)
|
||||
|
||||
|
||||
# Additional tests for tooling coverage
|
||||
|
||||
|
||||
@@ -176,11 +186,11 @@ def test_merge_tools_no_client_tools() -> None:
|
||||
|
||||
|
||||
def test_merge_tools_all_duplicates() -> None:
|
||||
"""merge_tools returns None when all client tools duplicate server tools."""
|
||||
"""merge_tools raises when client and server tools share a name."""
|
||||
server = [DummyTool("a"), DummyTool("b")]
|
||||
client = [DummyTool("a"), DummyTool("b")]
|
||||
result = merge_tools(server, client)
|
||||
assert result is None
|
||||
with pytest.raises(ValueError, match="Duplicate tool name 'a'"):
|
||||
merge_tools(server, client)
|
||||
|
||||
|
||||
def test_merge_tools_empty_server() -> None:
|
||||
@@ -208,7 +218,7 @@ def test_merge_tools_with_approval_tools_no_client() -> None:
|
||||
|
||||
|
||||
def test_merge_tools_with_approval_tools_all_duplicates() -> None:
|
||||
"""merge_tools returns server tools with approval mode even when client duplicates."""
|
||||
"""merge_tools raises even when a client tool duplicates an approval-gated server tool."""
|
||||
|
||||
class ApprovalTool:
|
||||
def __init__(self, name: str):
|
||||
@@ -217,7 +227,5 @@ def test_merge_tools_with_approval_tools_all_duplicates() -> None:
|
||||
|
||||
server = [ApprovalTool("write_doc")]
|
||||
client = [DummyTool("write_doc")] # Same name as server
|
||||
result = merge_tools(server, client)
|
||||
assert result is not None
|
||||
assert len(result) == 1
|
||||
assert result[0].approval_mode == "always_require"
|
||||
with pytest.raises(ValueError, match="Duplicate tool name 'write_doc'"):
|
||||
merge_tools(server, client)
|
||||
|
||||
@@ -228,11 +228,11 @@ class AnthropicClient(
|
||||
model_id: str | None = None,
|
||||
anthropic_client: AsyncAnthropic | None = None,
|
||||
additional_beta_flags: list[str] | None = None,
|
||||
additional_properties: dict[str, Any] | None = None,
|
||||
middleware: Sequence[ChatAndFunctionMiddlewareTypes] | None = None,
|
||||
function_invocation_configuration: FunctionInvocationConfiguration | None = None,
|
||||
env_file_path: str | None = None,
|
||||
env_file_encoding: str | None = None,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
"""Initialize an Anthropic Agent client.
|
||||
|
||||
@@ -244,11 +244,11 @@ class AnthropicClient(
|
||||
For instance if you need to set a different base_url for testing or private deployments.
|
||||
additional_beta_flags: Additional beta flags to enable on the client.
|
||||
Default flags are: "mcp-client-2025-04-04", "code-execution-2025-08-25".
|
||||
additional_properties: Additional properties stored on the client instance.
|
||||
middleware: Optional middleware to apply to the client.
|
||||
function_invocation_configuration: Optional function invocation configuration override.
|
||||
env_file_path: Path to environment file for loading settings.
|
||||
env_file_encoding: Encoding of the environment file.
|
||||
kwargs: Additional keyword arguments passed to the parent class.
|
||||
|
||||
Examples:
|
||||
.. code-block:: python
|
||||
@@ -319,9 +319,9 @@ class AnthropicClient(
|
||||
|
||||
# Initialize parent
|
||||
super().__init__(
|
||||
additional_properties=additional_properties,
|
||||
middleware=middleware,
|
||||
function_invocation_configuration=function_invocation_configuration,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
# Initialize instance variables
|
||||
|
||||
@@ -24,7 +24,7 @@ classifiers = [
|
||||
]
|
||||
dependencies = [
|
||||
"agent-framework-core>=1.0.0rc4",
|
||||
"anthropic>=0.70.0,<1",
|
||||
"anthropic>=0.80.0,<0.80.1",
|
||||
]
|
||||
|
||||
[tool.uv]
|
||||
@@ -87,7 +87,7 @@ include = "../../shared_tasks.toml"
|
||||
|
||||
[tool.poe.tasks]
|
||||
mypy = "mypy --config-file $POE_ROOT/pyproject.toml agent_framework_anthropic"
|
||||
test = "pytest -m \"not integration\" --cov=agent_framework_anthropic --cov-report=term-missing:skip-covered -n auto --dist worksteal tests"
|
||||
test = 'pytest -m "not integration" --cov=agent_framework_anthropic --cov-report=term-missing:skip-covered -n auto --dist worksteal tests'
|
||||
|
||||
[build-system]
|
||||
requires = ["flit-core >= 3.11,<4.0"]
|
||||
|
||||
@@ -24,7 +24,7 @@ classifiers = [
|
||||
]
|
||||
dependencies = [
|
||||
"agent-framework-core>=1.0.0rc4",
|
||||
"azure-search-documents==11.7.0b2",
|
||||
"azure-search-documents>=11.7.0b2,<11.7.0b3",
|
||||
]
|
||||
|
||||
[tool.uv]
|
||||
@@ -89,7 +89,7 @@ include = "../../shared_tasks.toml"
|
||||
|
||||
[tool.poe.tasks]
|
||||
mypy = "mypy --config-file $POE_ROOT/pyproject.toml agent_framework_azure_ai_search"
|
||||
test = "pytest -m \"not integration\" --cov=agent_framework_azure_ai_search --cov-report=term-missing:skip-covered tests"
|
||||
test = 'pytest -m "not integration" --cov=agent_framework_azure_ai_search --cov-report=term-missing:skip-covered tests'
|
||||
|
||||
[build-system]
|
||||
requires = ["flit-core >= 3.11,<4.0"]
|
||||
|
||||
@@ -17,10 +17,15 @@ from agent_framework_azure_ai_search._context_provider import AzureAISearchConte
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def clear_azure_search_environment(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
for key in tuple(os.environ):
|
||||
if key.startswith("AZURE_SEARCH_"):
|
||||
monkeypatch.delenv(key, raising=False)
|
||||
def clear_azure_search_env(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
"""Keep tests isolated from ambient Azure Search environment variables."""
|
||||
for key in (
|
||||
"AZURE_SEARCH_ENDPOINT",
|
||||
"AZURE_SEARCH_INDEX_NAME",
|
||||
"AZURE_SEARCH_KNOWLEDGE_BASE_NAME",
|
||||
"AZURE_SEARCH_API_KEY",
|
||||
):
|
||||
monkeypatch.delenv(key, raising=False)
|
||||
|
||||
|
||||
class MockSearchResults:
|
||||
|
||||
@@ -444,11 +444,11 @@ class AzureAIAgentClient(
|
||||
model_deployment_name: str | None = None,
|
||||
credential: AzureCredentialTypes | None = None,
|
||||
should_cleanup_agent: bool = True,
|
||||
additional_properties: dict[str, Any] | None = None,
|
||||
middleware: Sequence[ChatAndFunctionMiddlewareTypes] | None = None,
|
||||
function_invocation_configuration: FunctionInvocationConfiguration | None = None,
|
||||
env_file_path: str | None = None,
|
||||
env_file_encoding: str | None = None,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
"""Initialize an Azure AI Agent client.
|
||||
|
||||
@@ -471,11 +471,11 @@ class AzureAIAgentClient(
|
||||
should_cleanup_agent: Whether to cleanup (delete) agents created by this client when
|
||||
the client is closed or context is exited. Defaults to True. Only affects agents
|
||||
created by this client instance; existing agents passed via agent_id are never deleted.
|
||||
additional_properties: Additional properties stored on the client instance.
|
||||
middleware: Optional sequence of middlewares to include.
|
||||
function_invocation_configuration: Optional function invocation configuration.
|
||||
env_file_path: Path to environment file for loading settings.
|
||||
env_file_encoding: Encoding of the environment file.
|
||||
kwargs: Additional keyword arguments passed to the parent class.
|
||||
|
||||
Examples:
|
||||
.. code-block:: python
|
||||
@@ -548,9 +548,9 @@ class AzureAIAgentClient(
|
||||
|
||||
# Initialize parent
|
||||
super().__init__(
|
||||
additional_properties=additional_properties,
|
||||
middleware=middleware,
|
||||
function_invocation_configuration=function_invocation_configuration,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
# Initialize instance variables
|
||||
|
||||
@@ -119,9 +119,9 @@ class RawAzureAIClient(RawOpenAIResponsesClient[AzureAIClientOptionsT], Generic[
|
||||
credential: AzureCredentialTypes | None = None,
|
||||
use_latest_version: bool | None = None,
|
||||
allow_preview: bool | None = None,
|
||||
additional_properties: dict[str, Any] | None = None,
|
||||
env_file_path: str | None = None,
|
||||
env_file_encoding: str | None = None,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
"""Initialize a bare Azure AI client.
|
||||
|
||||
@@ -145,9 +145,9 @@ class RawAzureAIClient(RawOpenAIResponsesClient[AzureAIClientOptionsT], Generic[
|
||||
use_latest_version: Boolean flag that indicates whether to use latest agent version
|
||||
if it exists in the service.
|
||||
allow_preview: Enables preview opt-in on internally-created ``AIProjectClient``.
|
||||
additional_properties: Additional properties stored on the client instance.
|
||||
env_file_path: Path to environment file for loading settings.
|
||||
env_file_encoding: Encoding of the environment file.
|
||||
kwargs: Additional keyword arguments passed to the parent class.
|
||||
|
||||
Examples:
|
||||
.. code-block:: python
|
||||
@@ -217,7 +217,7 @@ class RawAzureAIClient(RawOpenAIResponsesClient[AzureAIClientOptionsT], Generic[
|
||||
|
||||
# Initialize parent
|
||||
super().__init__(
|
||||
**kwargs,
|
||||
additional_properties=additional_properties,
|
||||
)
|
||||
|
||||
# Initialize instance variables
|
||||
@@ -1243,11 +1243,11 @@ class AzureAIClient(
|
||||
credential: AzureCredentialTypes | None = None,
|
||||
use_latest_version: bool | None = None,
|
||||
allow_preview: bool | None = None,
|
||||
additional_properties: dict[str, Any] | None = None,
|
||||
middleware: Sequence[ChatAndFunctionMiddlewareTypes] | None = None,
|
||||
function_invocation_configuration: FunctionInvocationConfiguration | None = None,
|
||||
env_file_path: str | None = None,
|
||||
env_file_encoding: str | None = None,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
"""Initialize an Azure AI client with full layer support.
|
||||
|
||||
@@ -1268,11 +1268,11 @@ class AzureAIClient(
|
||||
use_latest_version: Boolean flag that indicates whether to use latest agent version
|
||||
if it exists in the service.
|
||||
allow_preview: Enables preview opt-in on internally-created ``AIProjectClient``
|
||||
additional_properties: Additional properties stored on the client instance.
|
||||
middleware: Optional sequence of chat middlewares to include.
|
||||
function_invocation_configuration: Optional function invocation configuration.
|
||||
env_file_path: Path to environment file for loading settings.
|
||||
env_file_encoding: Encoding of the environment file.
|
||||
kwargs: Additional keyword arguments passed to the parent class.
|
||||
|
||||
Examples:
|
||||
.. code-block:: python
|
||||
@@ -1319,9 +1319,9 @@ class AzureAIClient(
|
||||
credential=credential,
|
||||
use_latest_version=use_latest_version,
|
||||
allow_preview=allow_preview,
|
||||
additional_properties=additional_properties,
|
||||
middleware=middleware,
|
||||
function_invocation_configuration=function_invocation_configuration,
|
||||
env_file_path=env_file_path,
|
||||
env_file_encoding=env_file_encoding,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
@@ -124,9 +124,9 @@ class RawAzureAIInferenceEmbeddingClient(
|
||||
text_client: EmbeddingsClient | None = None,
|
||||
image_client: ImageEmbeddingsClient | None = None,
|
||||
credential: AzureKeyCredential | None = None,
|
||||
additional_properties: dict[str, Any] | None = None,
|
||||
env_file_path: str | None = None,
|
||||
env_file_encoding: str | None = None,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
"""Initialize a raw Azure AI Inference embedding client."""
|
||||
settings = load_settings(
|
||||
@@ -160,7 +160,7 @@ class RawAzureAIInferenceEmbeddingClient(
|
||||
credential=credential, # type: ignore[arg-type]
|
||||
)
|
||||
self._endpoint = resolved_endpoint
|
||||
super().__init__(**kwargs)
|
||||
super().__init__(additional_properties=additional_properties)
|
||||
|
||||
async def close(self) -> None:
|
||||
"""Close the underlying SDK clients and release resources."""
|
||||
@@ -376,9 +376,9 @@ class AzureAIInferenceEmbeddingClient(
|
||||
image_client: ImageEmbeddingsClient | None = None,
|
||||
credential: AzureKeyCredential | None = None,
|
||||
otel_provider_name: str | None = None,
|
||||
additional_properties: dict[str, Any] | None = None,
|
||||
env_file_path: str | None = None,
|
||||
env_file_encoding: str | None = None,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
"""Initialize an Azure AI Inference embedding client."""
|
||||
super().__init__(
|
||||
@@ -389,8 +389,8 @@ class AzureAIInferenceEmbeddingClient(
|
||||
text_client=text_client,
|
||||
image_client=image_client,
|
||||
credential=credential,
|
||||
additional_properties=additional_properties,
|
||||
otel_provider_name=otel_provider_name,
|
||||
env_file_path=env_file_path,
|
||||
env_file_encoding=env_file_encoding,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
@@ -24,9 +24,9 @@ classifiers = [
|
||||
]
|
||||
dependencies = [
|
||||
"agent-framework-core>=1.0.0rc4",
|
||||
"azure-ai-agents == 1.2.0b5",
|
||||
"azure-ai-inference>=1.0.0b9",
|
||||
"aiohttp",
|
||||
"azure-ai-agents>=1.2.0b5,<1.2.0b6",
|
||||
"azure-ai-inference>=1.0.0b9,<1.0.0b10",
|
||||
"aiohttp>=3.7.0,<4",
|
||||
]
|
||||
|
||||
[tool.uv]
|
||||
@@ -87,7 +87,7 @@ include = "../../shared_tasks.toml"
|
||||
|
||||
[tool.poe.tasks]
|
||||
mypy = "mypy --config-file $POE_ROOT/pyproject.toml agent_framework_azure_ai"
|
||||
test = "pytest -m \"not integration\" --cov=agent_framework_azure_ai --cov-report=term-missing:skip-covered tests"
|
||||
test = 'pytest -m "not integration" --cov=agent_framework_azure_ai --cov-report=term-missing:skip-covered tests'
|
||||
|
||||
[tool.poe.tasks.integration-tests]
|
||||
cmd = """
|
||||
|
||||
@@ -124,7 +124,13 @@ class CosmosHistoryProvider(BaseHistoryProvider):
|
||||
|
||||
self._database_client = self._cosmos_client.get_database_client(self.database_name)
|
||||
|
||||
async def get_messages(self, session_id: str | None, **kwargs: Any) -> list[Message]:
|
||||
async def get_messages(
|
||||
self,
|
||||
session_id: str | None,
|
||||
*,
|
||||
state: dict[str, Any] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> list[Message]:
|
||||
"""Retrieve stored messages for this session from Azure Cosmos DB."""
|
||||
await self._ensure_container_proxy()
|
||||
session_key = self._session_partition_key(session_id)
|
||||
@@ -157,7 +163,14 @@ class CosmosHistoryProvider(BaseHistoryProvider):
|
||||
|
||||
return messages
|
||||
|
||||
async def save_messages(self, session_id: str | None, messages: Sequence[Message], **kwargs: Any) -> None:
|
||||
async def save_messages(
|
||||
self,
|
||||
session_id: str | None,
|
||||
messages: Sequence[Message],
|
||||
*,
|
||||
state: dict[str, Any] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
"""Persist messages for this session to Azure Cosmos DB."""
|
||||
if not messages:
|
||||
return
|
||||
|
||||
@@ -24,7 +24,7 @@ classifiers = [
|
||||
]
|
||||
dependencies = [
|
||||
"agent-framework-core>=1.0.0rc4",
|
||||
"azure-cosmos>=4.9.0",
|
||||
"azure-cosmos>=4.3.0,<5",
|
||||
]
|
||||
|
||||
[tool.uv]
|
||||
|
||||
@@ -24,8 +24,8 @@ classifiers = [
|
||||
dependencies = [
|
||||
"agent-framework-core>=1.0.0rc4",
|
||||
"agent-framework-durabletask",
|
||||
"azure-functions",
|
||||
"azure-functions-durable",
|
||||
"azure-functions>=1.24.0,<2",
|
||||
"azure-functions-durable>=1.3.1,<2",
|
||||
]
|
||||
|
||||
[dependency-groups]
|
||||
@@ -93,7 +93,7 @@ include = "../../shared_tasks.toml"
|
||||
|
||||
[tool.poe.tasks]
|
||||
mypy = "mypy --config-file $POE_ROOT/pyproject.toml agent_framework_azurefunctions"
|
||||
test = "pytest -m \"not integration\" --cov=agent_framework_azurefunctions --cov-report=term-missing:skip-covered tests"
|
||||
test = 'pytest -m "not integration" --cov=agent_framework_azurefunctions --cov-report=term-missing:skip-covered tests'
|
||||
|
||||
[build-system]
|
||||
requires = ["flit-core >= 3.11,<4.0"]
|
||||
|
||||
@@ -236,11 +236,11 @@ class BedrockChatClient(
|
||||
session_token: str | None = None,
|
||||
client: BaseClient | None = None,
|
||||
boto3_session: Boto3Session | None = None,
|
||||
additional_properties: dict[str, Any] | None = None,
|
||||
middleware: Sequence[ChatAndFunctionMiddlewareTypes] | None = None,
|
||||
function_invocation_configuration: FunctionInvocationConfiguration | None = None,
|
||||
env_file_path: str | None = None,
|
||||
env_file_encoding: str | None = None,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
"""Create a Bedrock chat client and load AWS credentials.
|
||||
|
||||
@@ -252,11 +252,11 @@ class BedrockChatClient(
|
||||
session_token: Optional AWS session token for temporary credentials.
|
||||
client: Preconfigured Bedrock runtime client; when omitted a boto3 session is created.
|
||||
boto3_session: Custom boto3 session used to build the runtime client if provided.
|
||||
additional_properties: Additional properties stored on the client instance.
|
||||
middleware: Optional sequence of middlewares to include.
|
||||
function_invocation_configuration: Optional function invocation configuration
|
||||
env_file_path: Optional .env file path used by ``BedrockSettings`` to load defaults.
|
||||
env_file_encoding: Encoding for the optional .env file.
|
||||
kwargs: Additional arguments forwarded to ``BaseChatClient``.
|
||||
|
||||
Examples:
|
||||
.. code-block:: python
|
||||
@@ -303,9 +303,9 @@ class BedrockChatClient(
|
||||
)
|
||||
|
||||
super().__init__(
|
||||
additional_properties=additional_properties,
|
||||
middleware=middleware,
|
||||
function_invocation_configuration=function_invocation_configuration,
|
||||
**kwargs,
|
||||
)
|
||||
self.model_id = chat_model_id
|
||||
self.region = region
|
||||
|
||||
@@ -104,9 +104,9 @@ class RawBedrockEmbeddingClient(
|
||||
session_token: str | None = None,
|
||||
client: BaseClient | None = None,
|
||||
boto3_session: Boto3Session | None = None,
|
||||
additional_properties: dict[str, Any] | None = None,
|
||||
env_file_path: str | None = None,
|
||||
env_file_encoding: str | None = None,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
"""Initialize a raw Bedrock embedding client."""
|
||||
settings = load_settings(
|
||||
@@ -145,7 +145,7 @@ class RawBedrockEmbeddingClient(
|
||||
|
||||
self.model_id: str = settings["embedding_model_id"] # type: ignore[assignment] # pyright: ignore[reportTypedDictNotRequiredAccess]
|
||||
self.region = resolved_region
|
||||
super().__init__(**kwargs)
|
||||
super().__init__(additional_properties=additional_properties)
|
||||
|
||||
def service_url(self) -> str:
|
||||
"""Get the URL of the service."""
|
||||
@@ -274,9 +274,9 @@ class BedrockEmbeddingClient(
|
||||
client: BaseClient | None = None,
|
||||
boto3_session: Boto3Session | None = None,
|
||||
otel_provider_name: str | None = None,
|
||||
additional_properties: dict[str, Any] | None = None,
|
||||
env_file_path: str | None = None,
|
||||
env_file_encoding: str | None = None,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
"""Initialize a Bedrock embedding client."""
|
||||
super().__init__(
|
||||
@@ -287,8 +287,8 @@ class BedrockEmbeddingClient(
|
||||
session_token=session_token,
|
||||
client=client,
|
||||
boto3_session=boto3_session,
|
||||
additional_properties=additional_properties,
|
||||
otel_provider_name=otel_provider_name,
|
||||
env_file_path=env_file_path,
|
||||
env_file_encoding=env_file_encoding,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
@@ -86,7 +86,7 @@ include = "../../shared_tasks.toml"
|
||||
|
||||
[tool.poe.tasks]
|
||||
mypy = "mypy --config-file $POE_ROOT/pyproject.toml agent_framework_bedrock"
|
||||
test = "pytest -m \"not integration\" --cov=agent_framework_bedrock --cov-report=term-missing:skip-covered tests"
|
||||
test = 'pytest -m "not integration" --cov=agent_framework_bedrock --cov-report=term-missing:skip-covered tests'
|
||||
|
||||
[build-system]
|
||||
requires = ["hatchling"]
|
||||
|
||||
@@ -23,7 +23,7 @@ classifiers = [
|
||||
]
|
||||
dependencies = [
|
||||
"agent-framework-core>=1.0.0rc4",
|
||||
"openai-chatkit>=1.4.0,<2.0.0",
|
||||
"openai-chatkit>=1.4.1,<2.0.0",
|
||||
]
|
||||
|
||||
[tool.uv]
|
||||
@@ -88,7 +88,7 @@ include = "../../shared_tasks.toml"
|
||||
|
||||
[tool.poe.tasks]
|
||||
mypy = "mypy --config-file $POE_ROOT/pyproject.toml agent_framework_chatkit"
|
||||
test = "pytest -m \"not integration\" --cov=agent_framework_chatkit --cov-report=term-missing:skip-covered tests"
|
||||
test = 'pytest -m "not integration" --cov=agent_framework_chatkit --cov-report=term-missing:skip-covered tests'
|
||||
|
||||
[build-system]
|
||||
requires = ["flit-core >= 3.11,<4.0"]
|
||||
|
||||
@@ -590,6 +590,7 @@ class RawClaudeAgent(BaseAgent, Generic[OptionsT]):
|
||||
*,
|
||||
stream: Literal[False] = ...,
|
||||
session: AgentSession | None = None,
|
||||
options: OptionsT | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Awaitable[AgentResponse[Any]]: ...
|
||||
|
||||
@@ -600,6 +601,7 @@ class RawClaudeAgent(BaseAgent, Generic[OptionsT]):
|
||||
*,
|
||||
stream: Literal[True],
|
||||
session: AgentSession | None = None,
|
||||
options: OptionsT | None = None,
|
||||
**kwargs: Any,
|
||||
) -> ResponseStream[AgentResponseUpdate, AgentResponse[Any]]: ...
|
||||
|
||||
@@ -609,7 +611,8 @@ class RawClaudeAgent(BaseAgent, Generic[OptionsT]):
|
||||
*,
|
||||
stream: bool = False,
|
||||
session: AgentSession | None = None,
|
||||
**kwargs: Any,
|
||||
options: OptionsT | None = None,
|
||||
**kwargs: Any, # type: ignore
|
||||
) -> Awaitable[AgentResponse[Any]] | ResponseStream[AgentResponseUpdate, AgentResponse[Any]]:
|
||||
"""Run the agent with the given messages.
|
||||
|
||||
@@ -621,16 +624,16 @@ class RawClaudeAgent(BaseAgent, Generic[OptionsT]):
|
||||
returns an awaitable AgentResponse.
|
||||
session: The conversation session. If session has service_session_id set,
|
||||
the agent will resume that session.
|
||||
kwargs: Additional keyword arguments including 'options' for runtime options
|
||||
(model, permission_mode can be changed per-request).
|
||||
options: Runtime options. Model and permission_mode can be changed per request.
|
||||
kwargs: Additional keyword arguments for compatibility with the shared agent
|
||||
interface (e.g. compaction_strategy, tokenizer). Not used by ClaudeAgent.
|
||||
|
||||
Returns:
|
||||
When stream=True: An ResponseStream for streaming updates.
|
||||
When stream=False: An Awaitable[AgentResponse] with the complete response.
|
||||
"""
|
||||
options = kwargs.pop("options", None)
|
||||
response = ResponseStream(
|
||||
self._get_stream(messages, session=session, options=options, **kwargs),
|
||||
self._get_stream(messages, session=session, options=options),
|
||||
finalizer=self._finalize_response,
|
||||
)
|
||||
|
||||
@@ -643,8 +646,7 @@ class RawClaudeAgent(BaseAgent, Generic[OptionsT]):
|
||||
messages: AgentRunInputs | None = None,
|
||||
*,
|
||||
session: AgentSession | None = None,
|
||||
options: OptionsT | MutableMapping[str, Any] | None = None,
|
||||
**kwargs: Any,
|
||||
options: OptionsT | None = None,
|
||||
) -> AsyncIterable[AgentResponseUpdate]:
|
||||
"""Internal streaming implementation."""
|
||||
session = session or self.create_session()
|
||||
|
||||
@@ -24,7 +24,7 @@ classifiers = [
|
||||
]
|
||||
dependencies = [
|
||||
"agent-framework-core>=1.0.0rc4",
|
||||
"claude-agent-sdk>=0.1.25",
|
||||
"claude-agent-sdk>=0.1.36,<0.1.49",
|
||||
]
|
||||
|
||||
[tool.uv]
|
||||
@@ -88,7 +88,7 @@ include = "../../shared_tasks.toml"
|
||||
|
||||
[tool.poe.tasks]
|
||||
mypy = "mypy --config-file $POE_ROOT/pyproject.toml agent_framework_claude"
|
||||
test = "pytest -m \"not integration\" --cov=agent_framework_claude --cov-report=term-missing:skip-covered tests"
|
||||
test = 'pytest -m "not integration" --cov=agent_framework_claude --cov-report=term-missing:skip-covered tests'
|
||||
|
||||
[build-system]
|
||||
requires = ["flit-core >= 3.11,<4.0"]
|
||||
|
||||
@@ -196,7 +196,6 @@ class CopilotStudioAgent(BaseAgent):
|
||||
*,
|
||||
stream: Literal[False] = False,
|
||||
session: AgentSession | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Awaitable[AgentResponse]: ...
|
||||
|
||||
@overload
|
||||
@@ -206,7 +205,6 @@ class CopilotStudioAgent(BaseAgent):
|
||||
*,
|
||||
stream: Literal[True],
|
||||
session: AgentSession | None = None,
|
||||
**kwargs: Any,
|
||||
) -> ResponseStream[AgentResponseUpdate, AgentResponse]: ...
|
||||
|
||||
def run(
|
||||
@@ -215,7 +213,6 @@ class CopilotStudioAgent(BaseAgent):
|
||||
*,
|
||||
stream: bool = False,
|
||||
session: AgentSession | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Awaitable[AgentResponse] | ResponseStream[AgentResponseUpdate, AgentResponse]:
|
||||
"""Get a response from the agent.
|
||||
|
||||
@@ -229,22 +226,20 @@ class CopilotStudioAgent(BaseAgent):
|
||||
Keyword Args:
|
||||
stream: Whether to stream the response. Defaults to False.
|
||||
session: The conversation session associated with the message(s).
|
||||
kwargs: Additional keyword arguments.
|
||||
|
||||
Returns:
|
||||
When stream=False: An Awaitable[AgentResponse].
|
||||
When stream=True: A ResponseStream of AgentResponseUpdate items.
|
||||
"""
|
||||
if stream:
|
||||
return self._run_stream_impl(messages=messages, session=session, **kwargs)
|
||||
return self._run_impl(messages=messages, session=session, **kwargs)
|
||||
return self._run_stream_impl(messages=messages, session=session)
|
||||
return self._run_impl(messages=messages, session=session)
|
||||
|
||||
async def _run_impl(
|
||||
self,
|
||||
messages: AgentRunInputs | None = None,
|
||||
*,
|
||||
session: AgentSession | None = None,
|
||||
**kwargs: Any,
|
||||
) -> AgentResponse:
|
||||
"""Non-streaming implementation of run."""
|
||||
if not session:
|
||||
@@ -269,7 +264,6 @@ class CopilotStudioAgent(BaseAgent):
|
||||
messages: AgentRunInputs | None = None,
|
||||
*,
|
||||
session: AgentSession | None = None,
|
||||
**kwargs: Any,
|
||||
) -> ResponseStream[AgentResponseUpdate, AgentResponse]:
|
||||
"""Streaming implementation of run."""
|
||||
|
||||
|
||||
@@ -24,7 +24,7 @@ classifiers = [
|
||||
]
|
||||
dependencies = [
|
||||
"agent-framework-core>=1.0.0rc4",
|
||||
"microsoft-agents-copilotstudio-client>=0.3.1",
|
||||
"microsoft-agents-copilotstudio-client>=0.3.1,<0.3.2",
|
||||
]
|
||||
|
||||
[tool.uv]
|
||||
@@ -87,7 +87,7 @@ include = "../../shared_tasks.toml"
|
||||
|
||||
[tool.poe.tasks]
|
||||
mypy = "mypy --config-file $POE_ROOT/pyproject.toml agent_framework_copilotstudio"
|
||||
test = "pytest -m \"not integration\" --cov=agent_framework_copilotstudio --cov-report=term-missing:skip-covered tests"
|
||||
test = 'pytest -m "not integration" --cov=agent_framework_copilotstudio --cov-report=term-missing:skip-covered tests'
|
||||
|
||||
[build-system]
|
||||
requires = ["flit-core >= 3.11,<4.0"]
|
||||
|
||||
@@ -215,6 +215,7 @@ from ._workflows._workflow_executor import (
|
||||
)
|
||||
from .exceptions import (
|
||||
MiddlewareException,
|
||||
UserInputRequiredException,
|
||||
WorkflowCheckpointException,
|
||||
WorkflowConvergenceException,
|
||||
WorkflowException,
|
||||
@@ -349,6 +350,7 @@ __all__ = [
|
||||
"TypeCompatibilityError",
|
||||
"UpdateT",
|
||||
"UsageDetails",
|
||||
"UserInputRequiredException",
|
||||
"ValidationTypeEnum",
|
||||
"Workflow",
|
||||
"WorkflowAgent",
|
||||
|
||||
@@ -2,10 +2,10 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import inspect
|
||||
import logging
|
||||
import re
|
||||
import sys
|
||||
import warnings
|
||||
from collections.abc import Awaitable, Callable, Mapping, MutableMapping, Sequence
|
||||
from contextlib import AbstractAsyncContextManager, AsyncExitStack
|
||||
from copy import deepcopy
|
||||
@@ -27,11 +27,13 @@ from uuid import uuid4
|
||||
from mcp import types
|
||||
from mcp.server.lowlevel import Server
|
||||
from mcp.shared.exceptions import McpError
|
||||
from pydantic import BaseModel, Field, create_model
|
||||
from pydantic import BaseModel
|
||||
|
||||
from . import _tools as _tool_utils # pyright: ignore[reportPrivateUsage]
|
||||
from ._clients import BaseChatClient, SupportsChatGetResponse
|
||||
from ._docstrings import apply_layered_docstring
|
||||
from ._mcp import LOG_LEVEL_MAPPING, MCPTool
|
||||
from ._middleware import AgentMiddlewareLayer, MiddlewareTypes
|
||||
from ._middleware import AgentMiddlewareLayer, FunctionInvocationContext, MiddlewareTypes
|
||||
from ._serialization import SerializationMixin
|
||||
from ._sessions import (
|
||||
AgentSession,
|
||||
@@ -40,12 +42,7 @@ from ._sessions import (
|
||||
InMemoryHistoryProvider,
|
||||
SessionContext,
|
||||
)
|
||||
from ._tools import (
|
||||
FunctionInvocationLayer,
|
||||
FunctionTool,
|
||||
ToolTypes,
|
||||
normalize_tools,
|
||||
)
|
||||
from ._tools import FunctionInvocationLayer, FunctionTool, ToolTypes, normalize_tools
|
||||
from ._types import (
|
||||
AgentResponse,
|
||||
AgentResponseUpdate,
|
||||
@@ -57,7 +54,7 @@ from ._types import (
|
||||
map_chat_to_agent_update,
|
||||
normalize_messages,
|
||||
)
|
||||
from .exceptions import AgentInvalidResponseException
|
||||
from .exceptions import AgentInvalidResponseException, UserInputRequiredException
|
||||
from .observability import AgentTelemetryLayer
|
||||
|
||||
if sys.version_info >= (3, 13):
|
||||
@@ -79,6 +76,9 @@ if TYPE_CHECKING:
|
||||
|
||||
logger = logging.getLogger("agent_framework")
|
||||
|
||||
_append_unique_tools = _tool_utils._append_unique_tools # pyright: ignore[reportPrivateUsage]
|
||||
_get_tool_name = _tool_utils._get_tool_name # pyright: ignore[reportPrivateUsage]
|
||||
|
||||
ResponseModelBoundT = TypeVar("ResponseModelBoundT", bound=BaseModel)
|
||||
OptionsCoT = TypeVar(
|
||||
"OptionsCoT",
|
||||
@@ -88,19 +88,6 @@ OptionsCoT = TypeVar(
|
||||
)
|
||||
|
||||
|
||||
def _get_tool_name(tool: Any) -> str | None:
|
||||
"""Extract a tool's name from either an object with a .name attribute or a dict tool definition."""
|
||||
if isinstance(tool, Mapping):
|
||||
tool_mapping = cast(Mapping[str, Any], tool)
|
||||
func = tool_mapping.get("function")
|
||||
if isinstance(func, Mapping):
|
||||
func_mapping = cast(Mapping[str, Any], func)
|
||||
name = func_mapping.get("name")
|
||||
return name if isinstance(name, str) else None
|
||||
return None
|
||||
return getattr(tool, "name", None)
|
||||
|
||||
|
||||
def _merge_options(base: dict[str, Any], override: dict[str, Any]) -> dict[str, Any]:
|
||||
"""Merge two options dicts, with override values taking precedence.
|
||||
|
||||
@@ -115,11 +102,14 @@ def _merge_options(base: dict[str, Any], override: dict[str, Any]) -> dict[str,
|
||||
for key, value in override.items():
|
||||
if value is None:
|
||||
continue
|
||||
if key == "tools" and result.get("tools"):
|
||||
# Combine tool lists, avoiding duplicates by name
|
||||
existing_names = {_get_tool_name(t) for t in result["tools"]} - {None}
|
||||
unique_new = [t for t in value if _get_tool_name(t) not in existing_names]
|
||||
result["tools"] = list(result["tools"]) + unique_new
|
||||
if key == "tools" and (result.get("tools") or value):
|
||||
base_tools = normalize_tools(result.get("tools"))
|
||||
override_tools = normalize_tools(value)
|
||||
result["tools"] = _append_unique_tools(
|
||||
list(base_tools),
|
||||
override_tools,
|
||||
duplicate_error_message="Tool names must be unique.",
|
||||
)
|
||||
elif key == "logit_bias" and result.get("logit_bias"):
|
||||
# Merge logit_bias dicts
|
||||
result["logit_bias"] = {**result["logit_bias"], **value}
|
||||
@@ -180,8 +170,8 @@ class _RunContext(TypedDict):
|
||||
chat_options: MutableMapping[str, Any]
|
||||
compaction_strategy: CompactionStrategy | None
|
||||
tokenizer: TokenizerProtocol | None
|
||||
filtered_kwargs: Mapping[str, Any]
|
||||
finalize_kwargs: Mapping[str, Any]
|
||||
client_kwargs: Mapping[str, Any]
|
||||
function_invocation_kwargs: Mapping[str, Any]
|
||||
|
||||
|
||||
# region Agent Protocol
|
||||
@@ -229,15 +219,15 @@ class SupportsAgentRun(Protocol):
|
||||
|
||||
return AgentResponse(messages=[], response_id="custom-response")
|
||||
|
||||
def create_session(self, **kwargs):
|
||||
def create_session(self, *, session_id: str | None = None):
|
||||
from agent_framework import AgentSession
|
||||
|
||||
return AgentSession(**kwargs)
|
||||
return AgentSession(session_id=session_id)
|
||||
|
||||
def get_session(self, *, service_session_id, **kwargs):
|
||||
def get_session(self, service_session_id: str, *, session_id: str | None = None):
|
||||
from agent_framework import AgentSession
|
||||
|
||||
return AgentSession(service_session_id=service_session_id, **kwargs)
|
||||
return AgentSession(service_session_id=service_session_id, session_id=session_id)
|
||||
|
||||
|
||||
# Verify the instance satisfies the protocol
|
||||
@@ -256,6 +246,8 @@ class SupportsAgentRun(Protocol):
|
||||
*,
|
||||
stream: Literal[False] = ...,
|
||||
session: AgentSession | None = None,
|
||||
function_invocation_kwargs: Mapping[str, Any] | None = None,
|
||||
client_kwargs: Mapping[str, Any] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Awaitable[AgentResponse[Any]]:
|
||||
"""Get a response from the agent (non-streaming)."""
|
||||
@@ -268,6 +260,8 @@ class SupportsAgentRun(Protocol):
|
||||
*,
|
||||
stream: Literal[True],
|
||||
session: AgentSession | None = None,
|
||||
function_invocation_kwargs: Mapping[str, Any] | None = None,
|
||||
client_kwargs: Mapping[str, Any] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> ResponseStream[AgentResponseUpdate, AgentResponse[Any]]:
|
||||
"""Get a streaming response from the agent."""
|
||||
@@ -279,6 +273,8 @@ class SupportsAgentRun(Protocol):
|
||||
*,
|
||||
stream: bool = False,
|
||||
session: AgentSession | None = None,
|
||||
function_invocation_kwargs: Mapping[str, Any] | None = None,
|
||||
client_kwargs: Mapping[str, Any] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Awaitable[AgentResponse[Any]] | ResponseStream[AgentResponseUpdate, AgentResponse[Any]]:
|
||||
"""Get a response from the agent.
|
||||
@@ -293,6 +289,8 @@ class SupportsAgentRun(Protocol):
|
||||
Keyword Args:
|
||||
stream: Whether to stream the response. Defaults to False.
|
||||
session: The conversation session associated with the message(s).
|
||||
function_invocation_kwargs: Keyword arguments forwarded to tool invocation.
|
||||
client_kwargs: Additional client-specific keyword arguments.
|
||||
kwargs: Additional keyword arguments.
|
||||
|
||||
Returns:
|
||||
@@ -302,11 +300,11 @@ class SupportsAgentRun(Protocol):
|
||||
"""
|
||||
...
|
||||
|
||||
def create_session(self, **kwargs: Any) -> AgentSession:
|
||||
def create_session(self, *, session_id: str | None = None) -> AgentSession:
|
||||
"""Creates a new conversation session."""
|
||||
...
|
||||
|
||||
def get_session(self, *, service_session_id: str, **kwargs: Any) -> AgentSession:
|
||||
def get_session(self, service_session_id: str, *, session_id: str | None = None) -> AgentSession:
|
||||
"""Gets or creates a session for a service-managed session ID."""
|
||||
...
|
||||
|
||||
@@ -389,6 +387,13 @@ class BaseAgent(SerializationMixin):
|
||||
additional_properties: Additional properties set on the agent.
|
||||
kwargs: Additional keyword arguments (merged into additional_properties).
|
||||
"""
|
||||
if kwargs:
|
||||
warnings.warn(
|
||||
"Passing additional properties as direct keyword arguments to BaseAgent is deprecated; "
|
||||
"pass them via additional_properties instead.",
|
||||
DeprecationWarning,
|
||||
stacklevel=3,
|
||||
)
|
||||
if id is None:
|
||||
id = str(uuid4())
|
||||
self.id = id
|
||||
@@ -403,27 +408,40 @@ class BaseAgent(SerializationMixin):
|
||||
self.additional_properties: dict[str, Any] = cast(dict[str, Any], additional_properties or {})
|
||||
self.additional_properties.update(kwargs)
|
||||
|
||||
def create_session(self, *, session_id: str | None = None, **kwargs: Any) -> AgentSession:
|
||||
def create_session(self, *, session_id: str | None = None) -> AgentSession:
|
||||
"""Create a new lightweight session.
|
||||
|
||||
This will be used by an agent to hold the persisted session.
|
||||
This depends on the service used, in some cases, or with store=True
|
||||
this will add the ``service_session_id`` based on the response,
|
||||
which is then fed back to the API on the next call.
|
||||
|
||||
In other cases, if there is a HistoryProvider setup in the agent,
|
||||
that is used and it can store state in the session.
|
||||
|
||||
If there is no HistoryProvider and store=False or the default of a service is False.
|
||||
Then a ``InMemoryHistoryProvider`` instance is added to the agent and used with the session automatically.
|
||||
The ``InMemoryHistoryProvider`` stores the messages as `state` in the session by default.
|
||||
|
||||
Keyword Args:
|
||||
session_id: Optional session ID (generated if not provided).
|
||||
kwargs: Additional keyword arguments.
|
||||
|
||||
Returns:
|
||||
A new AgentSession instance.
|
||||
"""
|
||||
return AgentSession(session_id=session_id)
|
||||
|
||||
def get_session(self, *, service_session_id: str, session_id: str | None = None, **kwargs: Any) -> AgentSession:
|
||||
"""Get or create a session for a service-managed session ID.
|
||||
def get_session(self, service_session_id: str, *, session_id: str | None = None) -> AgentSession:
|
||||
"""Get a session for a service-managed session ID.
|
||||
|
||||
Only use this to create a session continuing that session id from a service.
|
||||
Otherwise use ``create_session``.
|
||||
|
||||
Args:
|
||||
service_session_id: The service-managed session ID.
|
||||
|
||||
Keyword Args:
|
||||
session_id: Optional local session ID (generated if not provided).
|
||||
kwargs: Additional keyword arguments.
|
||||
|
||||
Returns:
|
||||
A new AgentSession instance with service_session_id set.
|
||||
@@ -463,9 +481,8 @@ class BaseAgent(SerializationMixin):
|
||||
description: str | None = None,
|
||||
arg_name: str = "task",
|
||||
arg_description: str | None = None,
|
||||
stream_callback: Callable[[AgentResponseUpdate], None]
|
||||
| Callable[[AgentResponseUpdate], Awaitable[None]]
|
||||
| None = None,
|
||||
approval_mode: Literal["always_require", "never_require"] = "never_require",
|
||||
stream_callback: Callable[[AgentResponseUpdate], Awaitable[None] | None] | None = None,
|
||||
propagate_session: bool = False,
|
||||
) -> FunctionTool:
|
||||
"""Create a FunctionTool that wraps this agent.
|
||||
@@ -476,21 +493,15 @@ class BaseAgent(SerializationMixin):
|
||||
arg_name: The name of the function argument (default: "task").
|
||||
arg_description: The description for the function argument.
|
||||
If None, defaults to "Task for {tool_name}".
|
||||
approval_mode: Whether this delegated tool requires approval before execution.
|
||||
stream_callback: Optional callback for streaming responses. If provided, uses run(..., stream=True).
|
||||
propagate_session: If True, the parent agent's ``AgentSession`` is
|
||||
forwarded to this sub-agent's ``run()`` call, so both agents
|
||||
operate within the same logical session (sharing the same
|
||||
``session_id`` and provider-managed state, such as any stored
|
||||
conversation history or metadata). Defaults to False, meaning
|
||||
the sub-agent runs with a new, independent session.
|
||||
propagate_session: If True, the parent agent's session is forwarded
|
||||
to this sub-agent's ``run()`` call so both agents share the
|
||||
same session. Defaults to False.
|
||||
|
||||
Returns:
|
||||
A FunctionTool that can be used as a tool by other agents.
|
||||
|
||||
Raises:
|
||||
TypeError: If the agent does not implement SupportsAgentRun.
|
||||
ValueError: If the agent tool name cannot be determined.
|
||||
|
||||
Examples:
|
||||
.. code-block:: python
|
||||
|
||||
@@ -518,59 +529,46 @@ class BaseAgent(SerializationMixin):
|
||||
tool_description = description or self.description or ""
|
||||
argument_description = arg_description or f"Task for {tool_name}"
|
||||
|
||||
# Create dynamic input model with the specified argument name
|
||||
field_info = Field(..., description=argument_description)
|
||||
model_name = f"{name or _sanitize_agent_name(self.name) or 'agent'}_task"
|
||||
input_model = create_model(model_name, **{arg_name: (str, field_info)}) # type: ignore[call-overload]
|
||||
input_schema = {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
arg_name: {
|
||||
"type": "string",
|
||||
"description": argument_description,
|
||||
}
|
||||
},
|
||||
"required": [arg_name],
|
||||
"additionalProperties": False,
|
||||
}
|
||||
|
||||
# Check if callback is async once, outside the wrapper
|
||||
is_async_callback = stream_callback is not None and inspect.iscoroutinefunction(stream_callback)
|
||||
async def _agent_wrapper(ctx: FunctionInvocationContext, **kwargs: Any) -> str:
|
||||
"""Wrapper function that calls the agent.
|
||||
|
||||
async def agent_wrapper(**kwargs: Any) -> str:
|
||||
"""Wrapper function that calls the agent."""
|
||||
# Extract the input from kwargs using the specified arg_name
|
||||
input_text = kwargs.get(arg_name, "")
|
||||
Args:
|
||||
ctx: the function invocation context used
|
||||
**kwargs: only used to dynamically load the argument that is defined for this tool.
|
||||
"""
|
||||
stream = self.run(
|
||||
str(kwargs.get(arg_name, "")),
|
||||
stream=True,
|
||||
session=ctx.session if propagate_session else None,
|
||||
function_invocation_kwargs=dict(ctx.kwargs),
|
||||
)
|
||||
if stream_callback is not None:
|
||||
stream.with_transform_hook(stream_callback)
|
||||
final_response = await stream.get_final_response()
|
||||
if final_response.user_input_requests:
|
||||
raise UserInputRequiredException(contents=final_response.user_input_requests)
|
||||
# TODO(Copilot): update once #4331 merges
|
||||
return final_response.text
|
||||
|
||||
# Extract parent session when propagate_session is enabled
|
||||
parent_session = kwargs.get("session") if propagate_session else None
|
||||
|
||||
# Forward runtime context kwargs, excluding framework-internal keys.
|
||||
forwarded_kwargs = {
|
||||
k: v for k, v in kwargs.items() if k not in (arg_name, "conversation_id", "options", "session")
|
||||
}
|
||||
|
||||
if stream_callback is None:
|
||||
# Use non-streaming mode
|
||||
return (
|
||||
await self.run(
|
||||
input_text,
|
||||
stream=False,
|
||||
session=parent_session,
|
||||
**forwarded_kwargs,
|
||||
)
|
||||
).text
|
||||
|
||||
# Use streaming mode - accumulate updates and create final response
|
||||
response_updates: list[AgentResponseUpdate] = []
|
||||
async for update in self.run(input_text, stream=True, session=parent_session, **forwarded_kwargs):
|
||||
response_updates.append(update)
|
||||
if is_async_callback:
|
||||
await stream_callback(update) # type: ignore[misc]
|
||||
else:
|
||||
stream_callback(update)
|
||||
|
||||
# Create final text from accumulated updates
|
||||
return AgentResponse.from_updates(response_updates).text
|
||||
|
||||
agent_tool: FunctionTool = FunctionTool(
|
||||
return FunctionTool(
|
||||
name=tool_name,
|
||||
description=tool_description,
|
||||
func=agent_wrapper,
|
||||
input_model=input_model, # type: ignore
|
||||
approval_mode="never_require",
|
||||
func=_agent_wrapper,
|
||||
input_model=input_schema,
|
||||
approval_mode=approval_mode,
|
||||
)
|
||||
agent_tool._forward_runtime_kwargs = True # type: ignore
|
||||
return agent_tool
|
||||
|
||||
|
||||
# region Agent
|
||||
@@ -812,6 +810,8 @@ class RawAgent(BaseAgent, Generic[OptionsCoT]): # type: ignore[misc]
|
||||
options: ChatOptions[ResponseModelBoundT],
|
||||
compaction_strategy: CompactionStrategy | None = None,
|
||||
tokenizer: TokenizerProtocol | None = None,
|
||||
function_invocation_kwargs: Mapping[str, Any] | None = None,
|
||||
client_kwargs: Mapping[str, Any] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Awaitable[AgentResponse[ResponseModelBoundT]]: ...
|
||||
|
||||
@@ -826,6 +826,8 @@ class RawAgent(BaseAgent, Generic[OptionsCoT]): # type: ignore[misc]
|
||||
options: OptionsCoT | ChatOptions[None] | None = None,
|
||||
compaction_strategy: CompactionStrategy | None = None,
|
||||
tokenizer: TokenizerProtocol | None = None,
|
||||
function_invocation_kwargs: Mapping[str, Any] | None = None,
|
||||
client_kwargs: Mapping[str, Any] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Awaitable[AgentResponse[Any]]: ...
|
||||
|
||||
@@ -840,6 +842,8 @@ class RawAgent(BaseAgent, Generic[OptionsCoT]): # type: ignore[misc]
|
||||
options: OptionsCoT | ChatOptions[Any] | None = None,
|
||||
compaction_strategy: CompactionStrategy | None = None,
|
||||
tokenizer: TokenizerProtocol | None = None,
|
||||
function_invocation_kwargs: Mapping[str, Any] | None = None,
|
||||
client_kwargs: Mapping[str, Any] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> ResponseStream[AgentResponseUpdate, AgentResponse[Any]]: ...
|
||||
|
||||
@@ -853,6 +857,8 @@ class RawAgent(BaseAgent, Generic[OptionsCoT]): # type: ignore[misc]
|
||||
options: OptionsCoT | ChatOptions[Any] | None = None,
|
||||
compaction_strategy: CompactionStrategy | None = None,
|
||||
tokenizer: TokenizerProtocol | None = None,
|
||||
function_invocation_kwargs: Mapping[str, Any] | None = None,
|
||||
client_kwargs: Mapping[str, Any] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Awaitable[AgentResponse[Any]] | ResponseStream[AgentResponseUpdate, AgentResponse[Any]]:
|
||||
"""Run the agent with the given messages and options.
|
||||
@@ -882,14 +888,23 @@ class RawAgent(BaseAgent, Generic[OptionsCoT]): # type: ignore[misc]
|
||||
tokenizer: Optional per-run tokenizer override passed to
|
||||
``client.get_response()``. When omitted, the agent-level override
|
||||
is used, falling back to the client default.
|
||||
kwargs: Additional keyword arguments for the agent. These are only
|
||||
passed to functions that are called.
|
||||
function_invocation_kwargs: Keyword arguments forwarded to tool invocation.
|
||||
client_kwargs: Additional client-specific keyword arguments for the chat client.
|
||||
kwargs: Deprecated additional keyword arguments for the agent.
|
||||
They are forwarded to both tool invocation and the chat client for compatibility.
|
||||
|
||||
Returns:
|
||||
When stream=False: An Awaitable[AgentResponse] containing the agent's response.
|
||||
When stream=True: A ResponseStream of AgentResponseUpdate items with
|
||||
``get_final_response()`` for the final AgentResponse.
|
||||
"""
|
||||
if kwargs:
|
||||
warnings.warn(
|
||||
"Passing runtime keyword arguments directly to run() is deprecated; pass tool values via "
|
||||
"function_invocation_kwargs and client-specific values via client_kwargs instead.",
|
||||
DeprecationWarning,
|
||||
stacklevel=2,
|
||||
)
|
||||
if not stream:
|
||||
|
||||
async def _run_non_streaming() -> AgentResponse[Any]:
|
||||
@@ -900,7 +915,9 @@ class RawAgent(BaseAgent, Generic[OptionsCoT]): # type: ignore[misc]
|
||||
options=options,
|
||||
compaction_strategy=compaction_strategy,
|
||||
tokenizer=tokenizer,
|
||||
kwargs=kwargs,
|
||||
legacy_kwargs=kwargs,
|
||||
function_invocation_kwargs=function_invocation_kwargs,
|
||||
client_kwargs=client_kwargs,
|
||||
)
|
||||
response = cast(
|
||||
ChatResponse[Any],
|
||||
@@ -910,7 +927,8 @@ class RawAgent(BaseAgent, Generic[OptionsCoT]): # type: ignore[misc]
|
||||
options=ctx["chat_options"], # type: ignore[reportArgumentType]
|
||||
compaction_strategy=ctx["compaction_strategy"],
|
||||
tokenizer=ctx["tokenizer"],
|
||||
**ctx["filtered_kwargs"],
|
||||
function_invocation_kwargs=ctx["function_invocation_kwargs"],
|
||||
client_kwargs=ctx["client_kwargs"],
|
||||
),
|
||||
)
|
||||
|
||||
@@ -985,7 +1003,9 @@ class RawAgent(BaseAgent, Generic[OptionsCoT]): # type: ignore[misc]
|
||||
options=options,
|
||||
compaction_strategy=compaction_strategy,
|
||||
tokenizer=tokenizer,
|
||||
kwargs=kwargs,
|
||||
legacy_kwargs=kwargs,
|
||||
function_invocation_kwargs=function_invocation_kwargs,
|
||||
client_kwargs=client_kwargs,
|
||||
)
|
||||
ctx: _RunContext = ctx_holder["ctx"] # type: ignore[assignment] # Safe: we just assigned it
|
||||
return self.client.get_response( # type: ignore[call-overload, no-any-return]
|
||||
@@ -994,7 +1014,8 @@ class RawAgent(BaseAgent, Generic[OptionsCoT]): # type: ignore[misc]
|
||||
options=ctx["chat_options"], # type: ignore[reportArgumentType]
|
||||
compaction_strategy=ctx["compaction_strategy"],
|
||||
tokenizer=ctx["tokenizer"],
|
||||
**ctx["filtered_kwargs"],
|
||||
function_invocation_kwargs=ctx["function_invocation_kwargs"],
|
||||
client_kwargs=ctx["client_kwargs"],
|
||||
)
|
||||
|
||||
def _propagate_conversation_id(
|
||||
@@ -1082,9 +1103,12 @@ class RawAgent(BaseAgent, Generic[OptionsCoT]): # type: ignore[misc]
|
||||
options: Mapping[str, Any] | None,
|
||||
compaction_strategy: CompactionStrategy | None,
|
||||
tokenizer: TokenizerProtocol | None,
|
||||
kwargs: dict[str, Any],
|
||||
legacy_kwargs: Mapping[str, Any],
|
||||
function_invocation_kwargs: Mapping[str, Any] | None,
|
||||
client_kwargs: Mapping[str, Any] | None,
|
||||
) -> _RunContext:
|
||||
opts = dict(options) if options else {}
|
||||
existing_additional_args: dict[str, Any] = opts.pop("additional_function_arguments", None) or {}
|
||||
|
||||
# Get tools from options or named parameter (named param takes precedence)
|
||||
tools_ = tools if tools is not None else opts.pop("tools", None)
|
||||
@@ -1115,35 +1139,50 @@ class RawAgent(BaseAgent, Generic[OptionsCoT]): # type: ignore[misc]
|
||||
input_messages=input_messages,
|
||||
options=opts,
|
||||
)
|
||||
default_additional_args = chat_options.pop("additional_function_arguments", None)
|
||||
if isinstance(default_additional_args, Mapping):
|
||||
existing_additional_args = {
|
||||
**dict(cast(Mapping[str, Any], default_additional_args)),
|
||||
**existing_additional_args,
|
||||
}
|
||||
|
||||
agent_name = self._get_agent_name()
|
||||
base_tools = normalize_tools(chat_options.pop("tools", None))
|
||||
mcp_duplicate_message = "Tool names must be unique. Consider setting `tool_name_prefix` on the MCPTool."
|
||||
|
||||
# Normalize tools
|
||||
normalized_tools = normalize_tools(tools_)
|
||||
|
||||
# Resolve final tool list (runtime provided tools + local MCP server tools)
|
||||
final_tools: list[FunctionTool | Callable[..., Any] | dict[str, Any] | Any] = []
|
||||
# Resolve final tool list (configured tools + runtime provided tools + local MCP server tools)
|
||||
final_tools = list(base_tools)
|
||||
for tool in normalized_tools:
|
||||
if isinstance(tool, MCPTool):
|
||||
if not tool.is_connected:
|
||||
await self._async_exit_stack.enter_async_context(tool)
|
||||
final_tools.extend(tool.functions) # type: ignore
|
||||
_append_unique_tools(
|
||||
final_tools,
|
||||
tool.functions,
|
||||
duplicate_error_message=mcp_duplicate_message,
|
||||
)
|
||||
else:
|
||||
final_tools.append(tool) # type: ignore
|
||||
_append_unique_tools(final_tools, [tool]) # type: ignore[list-item]
|
||||
|
||||
existing_names = {name for t in final_tools if (name := _get_tool_name(t)) is not None}
|
||||
for mcp_server in self.mcp_tools:
|
||||
if not mcp_server.is_connected:
|
||||
await self._async_exit_stack.enter_async_context(mcp_server)
|
||||
final_tools.extend(f for f in mcp_server.functions if f.name not in existing_names)
|
||||
_append_unique_tools(
|
||||
final_tools,
|
||||
mcp_server.functions,
|
||||
duplicate_error_message=mcp_duplicate_message,
|
||||
)
|
||||
|
||||
# Merge runtime kwargs into additional_function_arguments so they're available
|
||||
# in function middleware context and tool invocation.
|
||||
existing_additional_args: dict[str, Any] = opts.pop("additional_function_arguments", None) or {}
|
||||
additional_function_arguments = {**kwargs, **existing_additional_args}
|
||||
# Include session so as_tool() wrappers with propagate_session=True can access it.
|
||||
if active_session is not None:
|
||||
additional_function_arguments["session"] = active_session
|
||||
# TODO(Copilot): Delete once direct ``run(**kwargs)`` compatibility is removed.
|
||||
# Legacy compatibility still fans out direct run kwargs into tool runtime kwargs.
|
||||
effective_function_invocation_kwargs = {
|
||||
**dict(legacy_kwargs),
|
||||
**(dict(function_invocation_kwargs) if function_invocation_kwargs is not None else {}),
|
||||
}
|
||||
additional_function_arguments = {**effective_function_invocation_kwargs, **existing_additional_args}
|
||||
|
||||
# Build options dict from run() options merged with provided options
|
||||
run_opts: dict[str, Any] = {
|
||||
@@ -1152,7 +1191,6 @@ class RawAgent(BaseAgent, Generic[OptionsCoT]): # type: ignore[misc]
|
||||
if active_session
|
||||
else opts.pop("conversation_id", None),
|
||||
"allow_multiple_tool_calls": opts.pop("allow_multiple_tool_calls", None),
|
||||
"additional_function_arguments": additional_function_arguments or None,
|
||||
"frequency_penalty": opts.pop("frequency_penalty", None),
|
||||
"logit_bias": opts.pop("logit_bias", None),
|
||||
"max_tokens": opts.pop("max_tokens", None),
|
||||
@@ -1164,7 +1202,7 @@ class RawAgent(BaseAgent, Generic[OptionsCoT]): # type: ignore[misc]
|
||||
"store": opts.pop("store", None),
|
||||
"temperature": opts.pop("temperature", None),
|
||||
"tool_choice": opts.pop("tool_choice", None),
|
||||
"tools": final_tools,
|
||||
"tools": final_tools or None,
|
||||
"top_p": opts.pop("top_p", None),
|
||||
"user": opts.pop("user", None),
|
||||
**opts, # Remaining options are provider-specific
|
||||
@@ -1176,11 +1214,14 @@ class RawAgent(BaseAgent, Generic[OptionsCoT]): # type: ignore[misc]
|
||||
# Build session_messages from session context: context messages + input messages
|
||||
session_messages: list[Message] = session_context.get_messages(include_input=True)
|
||||
|
||||
# Ensure session is forwarded in kwargs for tool invocation
|
||||
finalize_kwargs = dict(kwargs)
|
||||
finalize_kwargs["session"] = active_session
|
||||
# Filter chat_options from kwargs to prevent duplicate keyword argument
|
||||
filtered_kwargs = {k: v for k, v in finalize_kwargs.items() if k != "chat_options"}
|
||||
# TODO(Copilot): Delete once direct ``run(**kwargs)`` compatibility is removed.
|
||||
# Legacy compatibility still fans out direct run kwargs into client kwargs.
|
||||
effective_client_kwargs = {
|
||||
**dict(legacy_kwargs),
|
||||
**(dict(client_kwargs) if client_kwargs is not None else {}),
|
||||
}
|
||||
if active_session is not None:
|
||||
effective_client_kwargs["session"] = active_session
|
||||
|
||||
return {
|
||||
"session": active_session,
|
||||
@@ -1191,8 +1232,8 @@ class RawAgent(BaseAgent, Generic[OptionsCoT]): # type: ignore[misc]
|
||||
"chat_options": co,
|
||||
"compaction_strategy": compaction_strategy or self.compaction_strategy,
|
||||
"tokenizer": tokenizer or self.tokenizer,
|
||||
"filtered_kwargs": filtered_kwargs,
|
||||
"finalize_kwargs": finalize_kwargs,
|
||||
"client_kwargs": effective_client_kwargs,
|
||||
"function_invocation_kwargs": additional_function_arguments,
|
||||
}
|
||||
|
||||
async def _finalize_response(
|
||||
@@ -1442,6 +1483,58 @@ class Agent(
|
||||
For a minimal implementation without these features, use :class:`RawAgent`.
|
||||
"""
|
||||
|
||||
@overload
|
||||
def run(
|
||||
self,
|
||||
messages: AgentRunInputs | None = None,
|
||||
*,
|
||||
stream: Literal[False] = ...,
|
||||
session: AgentSession | None = None,
|
||||
function_invocation_kwargs: Mapping[str, Any] | None = None,
|
||||
client_kwargs: Mapping[str, Any] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Awaitable[AgentResponse[Any]]: ...
|
||||
|
||||
@overload
|
||||
def run(
|
||||
self,
|
||||
messages: AgentRunInputs | None = None,
|
||||
*,
|
||||
stream: Literal[True],
|
||||
session: AgentSession | None = None,
|
||||
function_invocation_kwargs: Mapping[str, Any] | None = None,
|
||||
client_kwargs: Mapping[str, Any] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> ResponseStream[AgentResponseUpdate, AgentResponse[Any]]: ...
|
||||
|
||||
def run(
|
||||
self,
|
||||
messages: AgentRunInputs | None = None,
|
||||
*,
|
||||
stream: bool = False,
|
||||
session: AgentSession | None = None,
|
||||
middleware: Sequence[MiddlewareTypes] | None = None,
|
||||
options: OptionsCoT | ChatOptions[Any] | None = None,
|
||||
function_invocation_kwargs: Mapping[str, Any] | None = None,
|
||||
client_kwargs: Mapping[str, Any] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Awaitable[AgentResponse[Any]] | ResponseStream[AgentResponseUpdate, AgentResponse[Any]]:
|
||||
"""Run the agent."""
|
||||
super_run = cast(
|
||||
"Callable[..., Awaitable[AgentResponse[Any]] | ResponseStream[AgentResponseUpdate, AgentResponse[Any]]]",
|
||||
super().run, # type: ignore[misc]
|
||||
)
|
||||
return super_run( # type: ignore[no-any-return]
|
||||
messages=messages,
|
||||
stream=stream,
|
||||
session=session,
|
||||
middleware=middleware,
|
||||
options=options,
|
||||
function_invocation_kwargs=function_invocation_kwargs,
|
||||
client_kwargs=client_kwargs,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
client: SupportsChatGetResponse[OptionsCoT],
|
||||
@@ -1473,3 +1566,34 @@ class Agent(
|
||||
tokenizer=tokenizer,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
|
||||
def _apply_agent_docstrings() -> None:
|
||||
"""Align public agent docstrings with the raw implementation."""
|
||||
apply_layered_docstring(
|
||||
AgentMiddlewareLayer.run,
|
||||
RawAgent.run,
|
||||
extra_keyword_args={
|
||||
"middleware": """
|
||||
Optional per-run agent, chat, and function middleware.
|
||||
Agent middleware wraps the run itself, while chat and function middleware are forwarded to the
|
||||
underlying chat-client stack for this call.
|
||||
""",
|
||||
},
|
||||
)
|
||||
apply_layered_docstring(AgentTelemetryLayer.run, AgentMiddlewareLayer.run)
|
||||
apply_layered_docstring(
|
||||
Agent.run,
|
||||
RawAgent.run,
|
||||
extra_keyword_args={
|
||||
"middleware": """
|
||||
Optional per-run agent, chat, and function middleware.
|
||||
Agent middleware wraps the run itself, while chat and function middleware are forwarded to the
|
||||
underlying chat-client stack for this call.
|
||||
""",
|
||||
},
|
||||
)
|
||||
apply_layered_docstring(Agent.__init__, RawAgent.__init__)
|
||||
|
||||
|
||||
_apply_agent_docstrings()
|
||||
|
||||
@@ -4,6 +4,7 @@ from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import sys
|
||||
import warnings
|
||||
from abc import ABC, abstractmethod
|
||||
from collections.abc import (
|
||||
AsyncIterable,
|
||||
@@ -27,6 +28,7 @@ from typing import (
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
from ._docstrings import apply_layered_docstring
|
||||
from ._serialization import SerializationMixin
|
||||
from ._tools import (
|
||||
FunctionInvocationConfiguration,
|
||||
@@ -105,7 +107,7 @@ class SupportsChatGetResponse(Protocol[OptionsContraT]):
|
||||
class CustomChatClient:
|
||||
additional_properties: dict = {}
|
||||
|
||||
def get_response(self, messages, *, stream=False, **kwargs):
|
||||
def get_response(self, messages, *, stream=False, client_kwargs=None, **kwargs):
|
||||
if stream:
|
||||
from agent_framework import ChatResponseUpdate, ResponseStream
|
||||
|
||||
@@ -149,6 +151,8 @@ class SupportsChatGetResponse(Protocol[OptionsContraT]):
|
||||
options: OptionsContraT | ChatOptions[None] | None = None,
|
||||
compaction_strategy: CompactionStrategy | None = None,
|
||||
tokenizer: TokenizerProtocol | None = None,
|
||||
function_invocation_kwargs: Mapping[str, Any] | None = None,
|
||||
client_kwargs: Mapping[str, Any] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Awaitable[ChatResponse[Any]]: ...
|
||||
|
||||
@@ -161,6 +165,8 @@ class SupportsChatGetResponse(Protocol[OptionsContraT]):
|
||||
options: OptionsContraT | ChatOptions[Any] | None = None,
|
||||
compaction_strategy: CompactionStrategy | None = None,
|
||||
tokenizer: TokenizerProtocol | None = None,
|
||||
function_invocation_kwargs: Mapping[str, Any] | None = None,
|
||||
client_kwargs: Mapping[str, Any] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> ResponseStream[ChatResponseUpdate, ChatResponse[Any]]: ...
|
||||
|
||||
@@ -172,6 +178,8 @@ class SupportsChatGetResponse(Protocol[OptionsContraT]):
|
||||
options: OptionsContraT | ChatOptions[Any] | None = None,
|
||||
compaction_strategy: CompactionStrategy | None = None,
|
||||
tokenizer: TokenizerProtocol | None = None,
|
||||
function_invocation_kwargs: Mapping[str, Any] | None = None,
|
||||
client_kwargs: Mapping[str, Any] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Awaitable[ChatResponse[Any]] | ResponseStream[ChatResponseUpdate, ChatResponse[Any]]:
|
||||
"""Send input and return the response.
|
||||
@@ -182,7 +190,9 @@ class SupportsChatGetResponse(Protocol[OptionsContraT]):
|
||||
options: Chat options as a TypedDict.
|
||||
compaction_strategy: Optional per-call compaction override.
|
||||
tokenizer: Optional per-call tokenizer override.
|
||||
**kwargs: Additional chat options.
|
||||
function_invocation_kwargs: Keyword arguments forwarded only to tool invocation layers.
|
||||
client_kwargs: Additional client-specific keyword arguments.
|
||||
**kwargs: Deprecated additional client-specific keyword arguments.
|
||||
|
||||
Returns:
|
||||
When stream=False: An awaitable ChatResponse from the client.
|
||||
@@ -283,23 +293,31 @@ class BaseChatClient(SerializationMixin, ABC, Generic[OptionsCoT]):
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
additional_properties: dict[str, Any] | None = None,
|
||||
compaction_strategy: CompactionStrategy | None = None,
|
||||
tokenizer: TokenizerProtocol | None = None,
|
||||
additional_properties: dict[str, Any] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
"""Initialize a BaseChatClient instance.
|
||||
|
||||
Keyword Args:
|
||||
additional_properties: Additional properties for the client.
|
||||
compaction_strategy: Optional compaction strategy to apply before model calls.
|
||||
tokenizer: Optional tokenizer used by token-aware compaction strategies.
|
||||
kwargs: Additional keyword arguments (merged into additional_properties).
|
||||
additional_properties: Additional properties for the client.
|
||||
kwargs: Additional keyword arguments (merged into additional_properties for now).
|
||||
"""
|
||||
self.additional_properties = additional_properties or {}
|
||||
self.compaction_strategy = compaction_strategy
|
||||
self.tokenizer = tokenizer
|
||||
super().__init__(**kwargs)
|
||||
if kwargs:
|
||||
warnings.warn(
|
||||
"Passing additional properties as direct keyword arguments to BaseChatClient is deprecated; "
|
||||
"pass them via additional_properties instead.",
|
||||
DeprecationWarning,
|
||||
stacklevel=3,
|
||||
)
|
||||
self.additional_properties.update(kwargs)
|
||||
super().__init__()
|
||||
|
||||
def to_dict(self, *, exclude: set[str] | None = None, exclude_none: bool = True) -> dict[str, Any]:
|
||||
"""Convert the instance to a dictionary.
|
||||
@@ -486,7 +504,13 @@ class BaseChatClient(SerializationMixin, ABC, Generic[OptionsCoT]):
|
||||
When omitted, the client-level default is used.
|
||||
tokenizer: Optional per-call tokenizer override. When omitted, the
|
||||
client-level default is used.
|
||||
**kwargs: Other keyword arguments, can be used to pass function specific parameters.
|
||||
**kwargs: Additional compatibility keyword arguments. Lower chat-client layers do not
|
||||
consume ``function_invocation_kwargs`` directly; if present, it is ignored here
|
||||
because function invocation has already been handled by upper layers. If a
|
||||
``client_kwargs`` mapping is present, it is flattened into standard keyword
|
||||
arguments before forwarding to ``_inner_get_response()`` so client implementations
|
||||
can leverage those values, while implementations that ignore
|
||||
extra kwargs remain compatible.
|
||||
|
||||
Returns:
|
||||
When streaming a response stream of ChatResponseUpdates, otherwise an Awaitable ChatResponse.
|
||||
@@ -495,12 +519,21 @@ class BaseChatClient(SerializationMixin, ABC, Generic[OptionsCoT]):
|
||||
compaction_strategy=compaction_strategy,
|
||||
tokenizer=tokenizer,
|
||||
)
|
||||
compatibility_client_kwargs = kwargs.pop("client_kwargs", None)
|
||||
kwargs.pop("function_invocation_kwargs", None)
|
||||
merged_client_kwargs = (
|
||||
dict(cast(Mapping[str, Any], compatibility_client_kwargs))
|
||||
if isinstance(compatibility_client_kwargs, Mapping)
|
||||
else {}
|
||||
)
|
||||
merged_client_kwargs.update(kwargs)
|
||||
|
||||
if not compaction_overrides:
|
||||
return self._inner_get_response(
|
||||
messages=messages,
|
||||
stream=stream,
|
||||
options=options or {},
|
||||
**kwargs,
|
||||
options=options or {}, # type: ignore[arg-type]
|
||||
**merged_client_kwargs,
|
||||
)
|
||||
|
||||
if stream:
|
||||
@@ -514,7 +547,7 @@ class BaseChatClient(SerializationMixin, ABC, Generic[OptionsCoT]):
|
||||
messages=prepared_messages,
|
||||
stream=True,
|
||||
options=options or {},
|
||||
**kwargs,
|
||||
**merged_client_kwargs,
|
||||
)
|
||||
if isinstance(stream_response, ResponseStream):
|
||||
return stream_response # type: ignore[reportUnknownVariableType]
|
||||
@@ -534,7 +567,7 @@ class BaseChatClient(SerializationMixin, ABC, Generic[OptionsCoT]):
|
||||
messages=prepared_messages,
|
||||
stream=False,
|
||||
options=options or {},
|
||||
**kwargs,
|
||||
**merged_client_kwargs,
|
||||
)
|
||||
|
||||
return _get_response()
|
||||
@@ -564,7 +597,7 @@ class BaseChatClient(SerializationMixin, ABC, Generic[OptionsCoT]):
|
||||
function_invocation_configuration: FunctionInvocationConfiguration | None = None,
|
||||
compaction_strategy: CompactionStrategy | None = None,
|
||||
tokenizer: TokenizerProtocol | None = None,
|
||||
**kwargs: Any,
|
||||
additional_properties: Mapping[str, Any] | None = None,
|
||||
) -> Agent[OptionsCoT]:
|
||||
"""Create a Agent with this client.
|
||||
|
||||
@@ -590,7 +623,7 @@ class BaseChatClient(SerializationMixin, ABC, Generic[OptionsCoT]):
|
||||
client-level compaction defaults remain in effect for each call.
|
||||
tokenizer: Optional agent-level tokenizer override. When omitted,
|
||||
client-level tokenizer defaults remain in effect for each call.
|
||||
kwargs: Any additional keyword arguments. Will be stored as ``additional_properties``.
|
||||
additional_properties: Additional properties stored on the created agent.
|
||||
|
||||
Returns:
|
||||
A Agent instance configured with this chat client.
|
||||
@@ -615,21 +648,24 @@ class BaseChatClient(SerializationMixin, ABC, Generic[OptionsCoT]):
|
||||
"""
|
||||
from ._agents import Agent
|
||||
|
||||
return Agent(
|
||||
client=self,
|
||||
id=id,
|
||||
name=name,
|
||||
description=description,
|
||||
instructions=instructions,
|
||||
tools=tools,
|
||||
default_options=cast(Any, default_options),
|
||||
context_providers=context_providers,
|
||||
middleware=middleware,
|
||||
function_invocation_configuration=function_invocation_configuration,
|
||||
compaction_strategy=compaction_strategy,
|
||||
tokenizer=tokenizer,
|
||||
**kwargs,
|
||||
)
|
||||
agent_kwargs: dict[str, Any] = {
|
||||
"client": self,
|
||||
"id": id,
|
||||
"name": name,
|
||||
"description": description,
|
||||
"instructions": instructions,
|
||||
"tools": tools,
|
||||
"default_options": cast(Any, default_options),
|
||||
"context_providers": context_providers,
|
||||
"middleware": middleware,
|
||||
"compaction_strategy": compaction_strategy,
|
||||
"tokenizer": tokenizer,
|
||||
"additional_properties": dict(additional_properties) if additional_properties is not None else None,
|
||||
}
|
||||
if function_invocation_configuration is not None:
|
||||
agent_kwargs["function_invocation_configuration"] = function_invocation_configuration
|
||||
|
||||
return Agent(**agent_kwargs)
|
||||
|
||||
|
||||
# endregion
|
||||
@@ -892,16 +928,14 @@ class BaseEmbeddingClient(SerializationMixin, ABC, Generic[EmbeddingInputT, Embe
|
||||
self,
|
||||
*,
|
||||
additional_properties: dict[str, Any] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
"""Initialize a BaseEmbeddingClient instance.
|
||||
|
||||
Args:
|
||||
additional_properties: Additional properties to pass to the client.
|
||||
**kwargs: Additional keyword arguments passed to parent classes (for MRO).
|
||||
"""
|
||||
self.additional_properties = additional_properties or {}
|
||||
super().__init__(**kwargs)
|
||||
super().__init__()
|
||||
|
||||
@abstractmethod
|
||||
async def get_embeddings(
|
||||
@@ -923,3 +957,36 @@ class BaseEmbeddingClient(SerializationMixin, ABC, Generic[EmbeddingInputT, Embe
|
||||
|
||||
|
||||
# endregion
|
||||
|
||||
|
||||
def _apply_get_response_docstrings() -> None:
|
||||
"""Align layered chat-client docstrings with the lowest public implementation."""
|
||||
from ._middleware import ChatMiddlewareLayer
|
||||
from ._tools import FunctionInvocationLayer
|
||||
from .observability import ChatTelemetryLayer
|
||||
|
||||
apply_layered_docstring(ChatTelemetryLayer.get_response, BaseChatClient.get_response)
|
||||
apply_layered_docstring(
|
||||
FunctionInvocationLayer.get_response,
|
||||
ChatTelemetryLayer.get_response,
|
||||
extra_keyword_args={
|
||||
"function_middleware": """
|
||||
Optional per-call function middleware.
|
||||
When omitted, middleware configured on the client or forwarded from higher layers is used.
|
||||
""",
|
||||
},
|
||||
)
|
||||
apply_layered_docstring(
|
||||
ChatMiddlewareLayer.get_response,
|
||||
FunctionInvocationLayer.get_response,
|
||||
extra_keyword_args={
|
||||
"middleware": """
|
||||
Optional per-call chat and function middleware.
|
||||
This compatibility keyword argument is merged with any ``client_kwargs["middleware"]`` value
|
||||
before the request is executed.
|
||||
""",
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
_apply_get_response_docstrings()
|
||||
|
||||
@@ -0,0 +1,85 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import inspect
|
||||
from collections.abc import Callable, Mapping
|
||||
from typing import Any
|
||||
|
||||
_GOOGLE_SECTION_HEADERS = (
|
||||
"Args:",
|
||||
"Keyword Args:",
|
||||
"Returns:",
|
||||
"Raises:",
|
||||
"Examples:",
|
||||
"Note:",
|
||||
"Notes:",
|
||||
"Warning:",
|
||||
"Warnings:",
|
||||
)
|
||||
|
||||
|
||||
def _find_section_index(lines: list[str], header: str) -> int | None:
|
||||
for index, line in enumerate(lines):
|
||||
if line == header:
|
||||
return index
|
||||
return None
|
||||
|
||||
|
||||
def _find_next_section_index(lines: list[str], start: int) -> int:
|
||||
for index in range(start, len(lines)):
|
||||
if lines[index] in _GOOGLE_SECTION_HEADERS:
|
||||
return index
|
||||
return len(lines)
|
||||
|
||||
|
||||
def _format_keyword_arg_lines(extra_keyword_args: Mapping[str, str]) -> list[str]:
|
||||
formatted_lines: list[str] = []
|
||||
for name, description in extra_keyword_args.items():
|
||||
description_lines = inspect.cleandoc(description).splitlines()
|
||||
if not description_lines:
|
||||
formatted_lines.append(f" {name}:")
|
||||
continue
|
||||
formatted_lines.append(f" {name}: {description_lines[0]}")
|
||||
formatted_lines.extend(f" {line}" for line in description_lines[1:])
|
||||
return formatted_lines
|
||||
|
||||
|
||||
def build_layered_docstring(
|
||||
source: Callable[..., Any],
|
||||
*,
|
||||
extra_keyword_args: Mapping[str, str] | None = None,
|
||||
) -> str | None:
|
||||
"""Build a Google-style docstring from a lower-layer implementation."""
|
||||
docstring = inspect.getdoc(source)
|
||||
if not docstring:
|
||||
return None
|
||||
if not extra_keyword_args:
|
||||
return docstring
|
||||
|
||||
lines = docstring.splitlines()
|
||||
formatted_keyword_arg_lines = _format_keyword_arg_lines(extra_keyword_args)
|
||||
keyword_args_index = _find_section_index(lines, "Keyword Args:")
|
||||
|
||||
if keyword_args_index is None:
|
||||
args_index = _find_section_index(lines, "Args:")
|
||||
if args_index is not None:
|
||||
insert_index = _find_next_section_index(lines, args_index + 1)
|
||||
else:
|
||||
insert_index = _find_next_section_index(lines, 0)
|
||||
lines[insert_index:insert_index] = ["", "Keyword Args:", *formatted_keyword_arg_lines]
|
||||
return "\n".join(lines).rstrip()
|
||||
|
||||
insert_index = _find_next_section_index(lines, keyword_args_index + 1)
|
||||
lines[insert_index:insert_index] = formatted_keyword_arg_lines
|
||||
return "\n".join(lines).rstrip()
|
||||
|
||||
|
||||
def apply_layered_docstring(
|
||||
target: Callable[..., Any],
|
||||
source: Callable[..., Any],
|
||||
*,
|
||||
extra_keyword_args: Mapping[str, str] | None = None,
|
||||
) -> None:
|
||||
"""Copy a lower-layer docstring onto a wrapper and extend it when needed."""
|
||||
target.__doc__ = build_layered_docstring(source, extra_keyword_args=extra_keyword_args)
|
||||
@@ -4,6 +4,7 @@ from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import base64
|
||||
import json
|
||||
import logging
|
||||
import re
|
||||
import sys
|
||||
@@ -26,9 +27,7 @@ from mcp.shared.exceptions import McpError
|
||||
from mcp.shared.session import RequestResponder
|
||||
from opentelemetry import propagate
|
||||
|
||||
from ._tools import (
|
||||
FunctionTool,
|
||||
)
|
||||
from ._tools import FunctionTool
|
||||
from ._types import (
|
||||
Content,
|
||||
Message,
|
||||
@@ -59,6 +58,8 @@ class MCPSpecificApproval(TypedDict, total=False):
|
||||
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
_MCP_REMOTE_NAME_KEY = "_mcp_remote_name"
|
||||
_MCP_NORMALIZED_NAME_KEY = "_mcp_normalized_name"
|
||||
|
||||
# region: Helpers
|
||||
|
||||
@@ -87,8 +88,6 @@ def _parse_prompt_result_from_mcp(
|
||||
Returns:
|
||||
A string representation of the prompt result.
|
||||
"""
|
||||
import json
|
||||
|
||||
parts: list[str] = []
|
||||
for message in mcp_type.messages:
|
||||
content = message.content
|
||||
@@ -194,7 +193,7 @@ def _parse_tool_result_from_mcp(
|
||||
result.append(Content.from_text(str(item)))
|
||||
|
||||
if not result:
|
||||
result.append(Content.from_text(""))
|
||||
result.append(Content.from_text("null"))
|
||||
return result
|
||||
|
||||
|
||||
@@ -372,6 +371,20 @@ def _normalize_mcp_name(name: str) -> str:
|
||||
return re.sub(r"[^A-Za-z0-9_.-]", "-", name)
|
||||
|
||||
|
||||
def _build_prefixed_mcp_name(
|
||||
normalized_name: str,
|
||||
tool_name_prefix: str | None,
|
||||
) -> str:
|
||||
"""Build the exposed MCP function name from a normalized name and optional prefix."""
|
||||
if not tool_name_prefix:
|
||||
return normalized_name
|
||||
normalized_prefix = _normalize_mcp_name(tool_name_prefix).rstrip("_.-")
|
||||
if not normalized_prefix:
|
||||
return normalized_name
|
||||
trimmed_name = normalized_name.lstrip("_.-")
|
||||
return f"{normalized_prefix}_{trimmed_name}" if trimmed_name else normalized_prefix
|
||||
|
||||
|
||||
def _inject_otel_into_mcp_meta(meta: dict[str, Any] | None = None) -> dict[str, Any] | None:
|
||||
"""Inject OpenTelemetry trace context into MCP request _meta via the global propagator(s)."""
|
||||
carrier: dict[str, str] = {}
|
||||
@@ -415,6 +428,7 @@ class MCPTool:
|
||||
description: str | None = None,
|
||||
approval_mode: (Literal["always_require", "never_require"] | MCPSpecificApproval | None) = None,
|
||||
allowed_tools: Collection[str] | None = None,
|
||||
tool_name_prefix: str | None = None,
|
||||
load_tools: bool = True,
|
||||
parse_tool_results: Callable[[types.CallToolResult], str | list[Content]] | None = None,
|
||||
load_prompts: bool = True,
|
||||
@@ -435,6 +449,7 @@ class MCPTool:
|
||||
description: A description of the MCP tool.
|
||||
approval_mode: Whether approval is required to run tools.
|
||||
allowed_tools: A collection of tool names to allow.
|
||||
tool_name_prefix: Optional prefix to prepend to exposed MCP function names.
|
||||
load_tools: Whether to load tools from the MCP server.
|
||||
parse_tool_results: An optional callable with signature
|
||||
``Callable[[types.CallToolResult], str]`` that overrides the default result
|
||||
@@ -458,12 +473,17 @@ class MCPTool:
|
||||
self.description = description or ""
|
||||
self.approval_mode = approval_mode
|
||||
self.allowed_tools = allowed_tools
|
||||
self.tool_name_prefix = _normalize_mcp_name(tool_name_prefix).rstrip("_.-") if tool_name_prefix else None
|
||||
self.additional_properties = additional_properties
|
||||
self.load_tools_flag = load_tools
|
||||
self.parse_tool_results = parse_tool_results
|
||||
self.load_prompts_flag = load_prompts
|
||||
self.parse_prompt_results = parse_prompt_results
|
||||
self._exit_stack = AsyncExitStack()
|
||||
self._lifecycle_lock = asyncio.Lock()
|
||||
self._lifecycle_request_lock = asyncio.Lock()
|
||||
self._lifecycle_queue: asyncio.Queue[tuple[str, bool, asyncio.Future[None]]] | None = None
|
||||
self._lifecycle_owner_task: asyncio.Task[None] | None = None
|
||||
self.session = session
|
||||
self.request_timeout = request_timeout
|
||||
self.client = client
|
||||
@@ -480,41 +500,127 @@ class MCPTool:
|
||||
"""Get the list of functions that are allowed."""
|
||||
if not self.allowed_tools:
|
||||
return self._functions
|
||||
return [func for func in self._functions if func.name in self.allowed_tools]
|
||||
allowed_names = set(self.allowed_tools)
|
||||
filtered_functions: list[FunctionTool] = []
|
||||
for func in self._functions:
|
||||
additional_properties = func.additional_properties or {}
|
||||
normalized_name = additional_properties.get(_MCP_NORMALIZED_NAME_KEY)
|
||||
remote_name = additional_properties.get(_MCP_REMOTE_NAME_KEY)
|
||||
if (
|
||||
func.name in allowed_names
|
||||
or (isinstance(normalized_name, str) and normalized_name in allowed_names)
|
||||
or (isinstance(remote_name, str) and remote_name in allowed_names)
|
||||
):
|
||||
filtered_functions.append(func)
|
||||
return filtered_functions
|
||||
|
||||
async def _ensure_lifecycle_owner(self) -> None:
|
||||
async with self._lifecycle_lock:
|
||||
if self._lifecycle_owner_task is not None and not self._lifecycle_owner_task.done():
|
||||
return
|
||||
|
||||
self._lifecycle_queue = asyncio.Queue()
|
||||
self._lifecycle_owner_task = asyncio.create_task(
|
||||
self._run_lifecycle_owner(),
|
||||
name=f"mcp-lifecycle:{self.name}",
|
||||
)
|
||||
|
||||
async def _run_lifecycle_owner(self) -> None:
|
||||
queue = self._lifecycle_queue
|
||||
if queue is None:
|
||||
return
|
||||
|
||||
stop_error: BaseException | None = None
|
||||
try:
|
||||
while True:
|
||||
action, reset, future = await queue.get()
|
||||
|
||||
try:
|
||||
if action == "connect":
|
||||
await self._connect_on_owner(reset=reset)
|
||||
elif action == "close":
|
||||
await self._close_on_owner()
|
||||
else:
|
||||
raise RuntimeError(f"Unknown MCP lifecycle action: {action}")
|
||||
except asyncio.CancelledError as ex:
|
||||
stop_error = ex
|
||||
if not future.done():
|
||||
future.set_exception(ex)
|
||||
raise
|
||||
except Exception as ex:
|
||||
if not future.done():
|
||||
future.set_exception(ex)
|
||||
else:
|
||||
if not future.done():
|
||||
future.set_result(None)
|
||||
|
||||
if action == "close":
|
||||
return
|
||||
except asyncio.CancelledError as ex:
|
||||
stop_error = ex
|
||||
raise
|
||||
finally:
|
||||
while True:
|
||||
try:
|
||||
_, _, future = queue.get_nowait()
|
||||
except asyncio.QueueEmpty:
|
||||
break
|
||||
if not future.done():
|
||||
future.set_exception(stop_error or RuntimeError("MCP lifecycle owner stopped unexpectedly."))
|
||||
|
||||
self._lifecycle_queue = None
|
||||
self._lifecycle_owner_task = None
|
||||
|
||||
def _is_lifecycle_owner_task(self) -> bool:
|
||||
owner_task = self._lifecycle_owner_task
|
||||
return owner_task is not None and asyncio.current_task() is owner_task
|
||||
|
||||
async def _run_on_lifecycle_owner(self, action: str, *, reset: bool = False) -> None:
|
||||
await self._ensure_lifecycle_owner()
|
||||
|
||||
if self._is_lifecycle_owner_task():
|
||||
if action == "connect":
|
||||
await self._connect_on_owner(reset=reset)
|
||||
elif action == "close":
|
||||
await self._close_on_owner()
|
||||
else:
|
||||
raise RuntimeError(f"Unknown MCP lifecycle action: {action}")
|
||||
return
|
||||
|
||||
queue = self._lifecycle_queue
|
||||
if queue is None:
|
||||
raise RuntimeError("MCP lifecycle owner is not available.")
|
||||
|
||||
future = asyncio.get_running_loop().create_future()
|
||||
await queue.put((action, reset, future))
|
||||
await future
|
||||
|
||||
async def _safe_close_exit_stack(self) -> None:
|
||||
"""Safely close the exit stack, handling cross-task boundary errors.
|
||||
|
||||
anyio's cancel scopes are bound to the task they were created in.
|
||||
If aclose() is called from a different task (e.g., during streaming reconnection),
|
||||
anyio will raise a RuntimeError or CancelledError. In this case, we log a warning
|
||||
and allow garbage collection to clean up the resources.
|
||||
|
||||
Known error variants:
|
||||
- "Attempted to exit cancel scope in a different task than it was entered in"
|
||||
- "Attempted to exit a cancel scope that isn't the current task's current cancel scope"
|
||||
- CancelledError from anyio cancel scope cleanup
|
||||
"""
|
||||
"""Safely close the exit stack, handling unexpected cleanup failures."""
|
||||
try:
|
||||
await self._exit_stack.aclose()
|
||||
except RuntimeError as e:
|
||||
error_msg = str(e).lower()
|
||||
# Check for anyio cancel scope errors (multiple variants exist)
|
||||
if "cancel scope" in error_msg:
|
||||
logger.warning(
|
||||
"Could not cleanly close MCP exit stack due to cancel scope error. "
|
||||
"Old resources will be garbage collected. Error: %s",
|
||||
"This indicates MCP lifecycle ownership was lost. Error: %s",
|
||||
e,
|
||||
)
|
||||
else:
|
||||
raise
|
||||
except asyncio.CancelledError:
|
||||
# CancelledError can occur during cleanup when cancel scopes are involved
|
||||
logger.warning(
|
||||
"Could not cleanly close MCP exit stack due to cancellation. Old resources will be garbage collected."
|
||||
)
|
||||
logger.warning("Could not cleanly close MCP exit stack because the lifecycle owner task was cancelled.")
|
||||
|
||||
async def connect(self, *, reset: bool = False) -> None:
|
||||
if self._is_lifecycle_owner_task():
|
||||
await self._connect_on_owner(reset=reset)
|
||||
return
|
||||
|
||||
async with self._lifecycle_request_lock:
|
||||
await self._run_on_lifecycle_owner("connect", reset=reset)
|
||||
|
||||
async def _connect_on_owner(self, *, reset: bool = False) -> None:
|
||||
"""Connect to the MCP server.
|
||||
|
||||
Establishes a connection to the MCP server, initializes the session,
|
||||
@@ -706,12 +812,16 @@ class MCPTool:
|
||||
|
||||
def _determine_approval_mode(
|
||||
self,
|
||||
local_name: str,
|
||||
*candidate_names: str,
|
||||
) -> Literal["always_require", "never_require"] | None:
|
||||
if isinstance(self.approval_mode, dict):
|
||||
if (always_require := self.approval_mode.get("always_require_approval")) and local_name in always_require:
|
||||
if (always_require := self.approval_mode.get("always_require_approval")) and any(
|
||||
name in always_require for name in candidate_names
|
||||
):
|
||||
return "always_require"
|
||||
if (never_require := self.approval_mode.get("never_require_approval")) and local_name in never_require:
|
||||
if (never_require := self.approval_mode.get("never_require_approval")) and any(
|
||||
name in never_require for name in candidate_names
|
||||
):
|
||||
return "never_require"
|
||||
return None
|
||||
return self.approval_mode # type: ignore[reportReturnType]
|
||||
@@ -736,20 +846,25 @@ class MCPTool:
|
||||
prompt_list = await self.session.list_prompts(params=params) # type: ignore[union-attr]
|
||||
|
||||
for prompt in prompt_list.prompts:
|
||||
local_name = _normalize_mcp_name(prompt.name)
|
||||
normalized_name = _normalize_mcp_name(prompt.name)
|
||||
local_name = _build_prefixed_mcp_name(normalized_name, self.tool_name_prefix)
|
||||
|
||||
# Skip if already loaded
|
||||
if local_name in existing_names:
|
||||
continue
|
||||
|
||||
input_model = _get_input_model_from_mcp_prompt(prompt)
|
||||
approval_mode = self._determine_approval_mode(local_name)
|
||||
approval_mode = self._determine_approval_mode(local_name, normalized_name, prompt.name)
|
||||
func: FunctionTool = FunctionTool(
|
||||
func=partial(self.get_prompt, prompt.name),
|
||||
name=local_name,
|
||||
description=prompt.description or "",
|
||||
approval_mode=approval_mode,
|
||||
input_model=input_model,
|
||||
additional_properties={
|
||||
_MCP_REMOTE_NAME_KEY: prompt.name,
|
||||
_MCP_NORMALIZED_NAME_KEY: normalized_name,
|
||||
},
|
||||
)
|
||||
self._functions.append(func)
|
||||
existing_names.add(local_name)
|
||||
@@ -779,13 +894,14 @@ class MCPTool:
|
||||
tool_list = await self.session.list_tools(params=params) # type: ignore[union-attr]
|
||||
|
||||
for tool in tool_list.tools:
|
||||
local_name = _normalize_mcp_name(tool.name)
|
||||
normalized_name = _normalize_mcp_name(tool.name)
|
||||
local_name = _build_prefixed_mcp_name(normalized_name, self.tool_name_prefix)
|
||||
|
||||
# Skip if already loaded
|
||||
if local_name in existing_names:
|
||||
continue
|
||||
|
||||
approval_mode = self._determine_approval_mode(local_name)
|
||||
approval_mode = self._determine_approval_mode(local_name, normalized_name, tool.name)
|
||||
# Create FunctionTools out of each tool
|
||||
func: FunctionTool = FunctionTool(
|
||||
func=partial(self.call_tool, tool.name),
|
||||
@@ -793,6 +909,10 @@ class MCPTool:
|
||||
description=tool.description or "",
|
||||
approval_mode=approval_mode,
|
||||
input_model=tool.inputSchema,
|
||||
additional_properties={
|
||||
_MCP_REMOTE_NAME_KEY: tool.name,
|
||||
_MCP_NORMALIZED_NAME_KEY: normalized_name,
|
||||
},
|
||||
)
|
||||
self._functions.append(func)
|
||||
existing_names.add(local_name)
|
||||
@@ -802,14 +922,23 @@ class MCPTool:
|
||||
break
|
||||
params = types.PaginatedRequestParams(cursor=tool_list.nextCursor)
|
||||
|
||||
async def _close_on_owner(self) -> None:
|
||||
await self._safe_close_exit_stack()
|
||||
self._exit_stack = AsyncExitStack()
|
||||
self.session = None
|
||||
self.is_connected = False
|
||||
|
||||
async def close(self) -> None:
|
||||
"""Disconnect from the MCP server.
|
||||
|
||||
Closes the connection and cleans up resources.
|
||||
"""
|
||||
await self._safe_close_exit_stack()
|
||||
self.session = None
|
||||
self.is_connected = False
|
||||
if self._is_lifecycle_owner_task():
|
||||
await self._close_on_owner()
|
||||
return
|
||||
|
||||
async with self._lifecycle_request_lock:
|
||||
await self._run_on_lifecycle_owner("close")
|
||||
|
||||
@abstractmethod
|
||||
def get_mcp_client(self) -> _AsyncGeneratorContextManager[Any, None]:
|
||||
@@ -1001,7 +1130,7 @@ class MCPTool:
|
||||
except ToolException:
|
||||
raise
|
||||
except Exception as ex:
|
||||
await self._safe_close_exit_stack()
|
||||
await self.close()
|
||||
raise ToolExecutionException("Failed to enter context manager.", inner_exception=ex) from ex
|
||||
|
||||
async def __aexit__(
|
||||
@@ -1055,6 +1184,7 @@ class MCPStdioTool(MCPTool):
|
||||
name: str,
|
||||
command: str,
|
||||
*,
|
||||
tool_name_prefix: str | None = None,
|
||||
load_tools: bool = True,
|
||||
parse_tool_results: Callable[[types.CallToolResult], str | list[Content]] | None = None,
|
||||
load_prompts: bool = True,
|
||||
@@ -1083,6 +1213,7 @@ class MCPStdioTool(MCPTool):
|
||||
command: The command to run the MCP server.
|
||||
|
||||
Keyword Args:
|
||||
tool_name_prefix: Optional prefix to prepend to exposed MCP function names.
|
||||
load_tools: Whether to load tools from the MCP server.
|
||||
parse_tool_results: An optional callable with signature
|
||||
``Callable[[types.CallToolResult], str]`` that overrides the default result
|
||||
@@ -1119,6 +1250,7 @@ class MCPStdioTool(MCPTool):
|
||||
description=description,
|
||||
approval_mode=approval_mode,
|
||||
allowed_tools=allowed_tools,
|
||||
tool_name_prefix=tool_name_prefix,
|
||||
additional_properties=additional_properties,
|
||||
session=session,
|
||||
client=client,
|
||||
@@ -1180,6 +1312,7 @@ class MCPStreamableHTTPTool(MCPTool):
|
||||
name: str,
|
||||
url: str,
|
||||
*,
|
||||
tool_name_prefix: str | None = None,
|
||||
load_tools: bool = True,
|
||||
parse_tool_results: Callable[[types.CallToolResult], str | list[Content]] | None = None,
|
||||
load_prompts: bool = True,
|
||||
@@ -1208,6 +1341,7 @@ class MCPStreamableHTTPTool(MCPTool):
|
||||
url: The URL of the MCP server.
|
||||
|
||||
Keyword Args:
|
||||
tool_name_prefix: Optional prefix to prepend to exposed MCP function names.
|
||||
load_tools: Whether to load tools from the MCP server.
|
||||
parse_tool_results: An optional callable with signature
|
||||
``Callable[[types.CallToolResult], str]`` that overrides the default result
|
||||
@@ -1246,6 +1380,7 @@ class MCPStreamableHTTPTool(MCPTool):
|
||||
description=description,
|
||||
approval_mode=approval_mode,
|
||||
allowed_tools=allowed_tools,
|
||||
tool_name_prefix=tool_name_prefix,
|
||||
additional_properties=additional_properties,
|
||||
session=session,
|
||||
client=client,
|
||||
@@ -1299,6 +1434,7 @@ class MCPWebsocketTool(MCPTool):
|
||||
name: str,
|
||||
url: str,
|
||||
*,
|
||||
tool_name_prefix: str | None = None,
|
||||
load_tools: bool = True,
|
||||
parse_tool_results: Callable[[types.CallToolResult], str | list[Content]] | None = None,
|
||||
load_prompts: bool = True,
|
||||
@@ -1325,6 +1461,7 @@ class MCPWebsocketTool(MCPTool):
|
||||
url: The URL of the MCP server.
|
||||
|
||||
Keyword Args:
|
||||
tool_name_prefix: Optional prefix to prepend to exposed MCP function names.
|
||||
load_tools: Whether to load tools from the MCP server.
|
||||
parse_tool_results: An optional callable with signature
|
||||
``Callable[[types.CallToolResult], str]`` that overrides the default result
|
||||
@@ -1358,6 +1495,7 @@ class MCPWebsocketTool(MCPTool):
|
||||
description=description,
|
||||
approval_mode=approval_mode,
|
||||
allowed_tools=allowed_tools,
|
||||
tool_name_prefix=tool_name_prefix,
|
||||
additional_properties=additional_properties,
|
||||
session=session,
|
||||
client=client,
|
||||
|
||||
@@ -109,7 +109,9 @@ class AgentContext:
|
||||
to see the actual execution result or can be set to override the execution result.
|
||||
For non-streaming: should be AgentResponse.
|
||||
For streaming: should be ResponseStream[AgentResponseUpdate, AgentResponse].
|
||||
kwargs: Additional keyword arguments passed to the agent run method.
|
||||
kwargs: Legacy runtime keyword arguments visible to agent middleware.
|
||||
client_kwargs: Client-specific keyword arguments for downstream chat clients.
|
||||
function_invocation_kwargs: Keyword arguments forwarded to tool invocation.
|
||||
|
||||
Examples:
|
||||
.. code-block:: python
|
||||
@@ -147,6 +149,8 @@ class AgentContext:
|
||||
metadata: Mapping[str, Any] | None = None,
|
||||
result: AgentResponse | ResponseStream[AgentResponseUpdate, AgentResponse] | None = None,
|
||||
kwargs: Mapping[str, Any] | None = None,
|
||||
client_kwargs: Mapping[str, Any] | None = None,
|
||||
function_invocation_kwargs: Mapping[str, Any] | None = None,
|
||||
stream_transform_hooks: Sequence[
|
||||
Callable[[AgentResponseUpdate], AgentResponseUpdate | Awaitable[AgentResponseUpdate]]
|
||||
]
|
||||
@@ -167,7 +171,9 @@ class AgentContext:
|
||||
tokenizer: Optional per-run tokenizer override.
|
||||
metadata: Metadata dictionary for sharing data between agent middleware.
|
||||
result: Agent execution result.
|
||||
kwargs: Additional keyword arguments passed to the agent run method.
|
||||
kwargs: Legacy runtime keyword arguments visible to agent middleware.
|
||||
client_kwargs: Client-specific keyword arguments for downstream chat clients.
|
||||
function_invocation_kwargs: Keyword arguments forwarded to tool invocation.
|
||||
stream_transform_hooks: Hooks to transform streamed updates.
|
||||
stream_result_hooks: Hooks to process the final result after streaming.
|
||||
stream_cleanup_hooks: Hooks to run after streaming completes.
|
||||
@@ -182,6 +188,10 @@ class AgentContext:
|
||||
self.metadata: dict[str, Any] = dict(metadata) if metadata is not None else {}
|
||||
self.result = result
|
||||
self.kwargs: dict[str, Any] = dict(kwargs) if kwargs is not None else {}
|
||||
self.client_kwargs: dict[str, Any] = dict(client_kwargs) if client_kwargs is not None else {}
|
||||
self.function_invocation_kwargs: dict[str, Any] = (
|
||||
dict(function_invocation_kwargs) if function_invocation_kwargs is not None else {}
|
||||
)
|
||||
self.stream_transform_hooks = list(stream_transform_hooks or [])
|
||||
self.stream_result_hooks = list(stream_result_hooks or [])
|
||||
self.stream_cleanup_hooks = list(stream_cleanup_hooks or [])
|
||||
@@ -196,11 +206,11 @@ class FunctionInvocationContext:
|
||||
Attributes:
|
||||
function: The function being invoked.
|
||||
arguments: The validated arguments for the function.
|
||||
session: The agent session for this invocation, if any.
|
||||
metadata: Metadata dictionary for sharing data between function middleware.
|
||||
result: Function execution result. Can be observed after calling ``call_next()``
|
||||
to see the actual execution result or can be set to override the execution result.
|
||||
|
||||
kwargs: Additional keyword arguments passed to the chat method that invoked this function.
|
||||
kwargs: Additional runtime keyword arguments forwarded to the function invocation.
|
||||
|
||||
Examples:
|
||||
.. code-block:: python
|
||||
@@ -225,6 +235,7 @@ class FunctionInvocationContext:
|
||||
self,
|
||||
function: FunctionTool,
|
||||
arguments: BaseModel | Mapping[str, Any],
|
||||
session: AgentSession | None = None,
|
||||
metadata: Mapping[str, Any] | None = None,
|
||||
result: Any = None,
|
||||
kwargs: Mapping[str, Any] | None = None,
|
||||
@@ -234,12 +245,14 @@ class FunctionInvocationContext:
|
||||
Args:
|
||||
function: The function being invoked.
|
||||
arguments: The validated arguments for the function.
|
||||
session: The agent session for this invocation, if any.
|
||||
metadata: Metadata dictionary for sharing data between function middleware.
|
||||
result: Function execution result.
|
||||
kwargs: Additional keyword arguments passed to the chat method that invoked this function.
|
||||
kwargs: Additional runtime keyword arguments forwarded to the function invocation.
|
||||
"""
|
||||
self.function = function
|
||||
self.arguments = arguments
|
||||
self.session = session
|
||||
self.metadata: dict[str, Any] = dict(metadata) if metadata is not None else {}
|
||||
self.result = result
|
||||
self.kwargs: dict[str, Any] = dict(kwargs) if kwargs is not None else {}
|
||||
@@ -262,6 +275,7 @@ class ChatContext:
|
||||
For non-streaming: should be ChatResponse.
|
||||
For streaming: should be ResponseStream[ChatResponseUpdate, ChatResponse].
|
||||
kwargs: Additional keyword arguments passed to the chat client.
|
||||
function_invocation_kwargs: Keyword arguments forwarded only to tool invocation layers.
|
||||
stream_transform_hooks: Hooks applied to transform each streamed update.
|
||||
stream_result_hooks: Hooks applied to the finalized response (after finalizer).
|
||||
stream_cleanup_hooks: Hooks executed after stream consumption (before finalizer).
|
||||
@@ -298,6 +312,7 @@ class ChatContext:
|
||||
metadata: Mapping[str, Any] | None = None,
|
||||
result: ChatResponse | ResponseStream[ChatResponseUpdate, ChatResponse] | None = None,
|
||||
kwargs: Mapping[str, Any] | None = None,
|
||||
function_invocation_kwargs: Mapping[str, Any] | None = None,
|
||||
stream_transform_hooks: Sequence[
|
||||
Callable[[ChatResponseUpdate], ChatResponseUpdate | Awaitable[ChatResponseUpdate]]
|
||||
]
|
||||
@@ -315,6 +330,7 @@ class ChatContext:
|
||||
metadata: Metadata dictionary for sharing data between chat middleware.
|
||||
result: Chat execution result.
|
||||
kwargs: Additional keyword arguments passed to the chat client.
|
||||
function_invocation_kwargs: Keyword arguments forwarded only to tool invocation layers.
|
||||
stream_transform_hooks: Transform hooks to apply to each streamed update.
|
||||
stream_result_hooks: Result hooks to apply to the finalized streaming response.
|
||||
stream_cleanup_hooks: Cleanup hooks to run after streaming completes.
|
||||
@@ -326,6 +342,9 @@ class ChatContext:
|
||||
self.metadata: dict[str, Any] = dict(metadata) if metadata is not None else {}
|
||||
self.result = result
|
||||
self.kwargs: dict[str, Any] = dict(kwargs) if kwargs is not None else {}
|
||||
self.function_invocation_kwargs: dict[str, Any] = (
|
||||
dict(function_invocation_kwargs) if function_invocation_kwargs is not None else {}
|
||||
)
|
||||
self.stream_transform_hooks = list(stream_transform_hooks or [])
|
||||
self.stream_result_hooks = list(stream_result_hooks or [])
|
||||
self.stream_cleanup_hooks = list(stream_cleanup_hooks or [])
|
||||
@@ -980,6 +999,7 @@ class ChatMiddlewareLayer(Generic[OptionsCoT]):
|
||||
options: ChatOptions[ResponseModelBoundT],
|
||||
compaction_strategy: CompactionStrategy | None = None,
|
||||
tokenizer: TokenizerProtocol | None = None,
|
||||
function_invocation_kwargs: Mapping[str, Any] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Awaitable[ChatResponse[ResponseModelBoundT]]: ...
|
||||
|
||||
@@ -992,6 +1012,8 @@ class ChatMiddlewareLayer(Generic[OptionsCoT]):
|
||||
options: OptionsCoT | ChatOptions[None] | None = None,
|
||||
compaction_strategy: CompactionStrategy | None = None,
|
||||
tokenizer: TokenizerProtocol | None = None,
|
||||
function_invocation_kwargs: Mapping[str, Any] | None = None,
|
||||
client_kwargs: Mapping[str, Any] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Awaitable[ChatResponse[Any]]: ...
|
||||
|
||||
@@ -1004,6 +1026,8 @@ class ChatMiddlewareLayer(Generic[OptionsCoT]):
|
||||
options: OptionsCoT | ChatOptions[Any] | None = None,
|
||||
compaction_strategy: CompactionStrategy | None = None,
|
||||
tokenizer: TokenizerProtocol | None = None,
|
||||
function_invocation_kwargs: Mapping[str, Any] | None = None,
|
||||
client_kwargs: Mapping[str, Any] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> ResponseStream[ChatResponseUpdate, ChatResponse[Any]]: ...
|
||||
|
||||
@@ -1015,6 +1039,8 @@ class ChatMiddlewareLayer(Generic[OptionsCoT]):
|
||||
options: OptionsCoT | ChatOptions[Any] | None = None,
|
||||
compaction_strategy: CompactionStrategy | None = None,
|
||||
tokenizer: TokenizerProtocol | None = None,
|
||||
function_invocation_kwargs: Mapping[str, Any] | None = None,
|
||||
client_kwargs: Mapping[str, Any] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Awaitable[ChatResponse[Any]] | ResponseStream[ChatResponseUpdate, ChatResponse[Any]]:
|
||||
"""Execute the chat pipeline if middleware is configured."""
|
||||
@@ -1025,9 +1051,10 @@ class ChatMiddlewareLayer(Generic[OptionsCoT]):
|
||||
if tokenizer is not None:
|
||||
kwargs["tokenizer"] = tokenizer
|
||||
|
||||
call_middleware = kwargs.pop("middleware", [])
|
||||
effective_client_kwargs = dict(client_kwargs) if client_kwargs is not None else {}
|
||||
call_middleware = kwargs.pop("middleware", effective_client_kwargs.pop("middleware", []))
|
||||
middleware = categorize_middleware(call_middleware)
|
||||
kwargs["function_middleware"] = middleware["function"]
|
||||
effective_client_kwargs["function_middleware"] = middleware["function"]
|
||||
|
||||
pipeline = ChatMiddlewarePipeline(
|
||||
*self.chat_middleware,
|
||||
@@ -1038,6 +1065,8 @@ class ChatMiddlewareLayer(Generic[OptionsCoT]):
|
||||
messages=messages,
|
||||
stream=stream,
|
||||
options=options,
|
||||
function_invocation_kwargs=function_invocation_kwargs,
|
||||
client_kwargs=effective_client_kwargs,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
@@ -1046,7 +1075,8 @@ class ChatMiddlewareLayer(Generic[OptionsCoT]):
|
||||
messages=list(messages),
|
||||
options=options,
|
||||
stream=stream,
|
||||
kwargs=kwargs,
|
||||
kwargs={**effective_client_kwargs, **kwargs},
|
||||
function_invocation_kwargs=function_invocation_kwargs,
|
||||
)
|
||||
|
||||
async def _execute() -> ChatResponse | ResponseStream[ChatResponseUpdate, ChatResponse] | None:
|
||||
@@ -1079,11 +1109,17 @@ class ChatMiddlewareLayer(Generic[OptionsCoT]):
|
||||
self, context: ChatContext
|
||||
) -> Awaitable[ChatResponse] | ResponseStream[ChatResponseUpdate, ChatResponse]:
|
||||
"""Internal middleware handler to adapt to pipeline."""
|
||||
handler_kwargs = dict(context.kwargs)
|
||||
compaction_strategy = handler_kwargs.pop("compaction_strategy", None)
|
||||
tokenizer = handler_kwargs.pop("tokenizer", None)
|
||||
return super().get_response( # type: ignore[misc, no-any-return]
|
||||
messages=context.messages,
|
||||
stream=context.stream,
|
||||
options=context.options or {},
|
||||
**context.kwargs,
|
||||
compaction_strategy=compaction_strategy,
|
||||
tokenizer=tokenizer,
|
||||
function_invocation_kwargs=context.function_invocation_kwargs,
|
||||
client_kwargs=handler_kwargs,
|
||||
)
|
||||
|
||||
|
||||
@@ -1115,6 +1151,8 @@ class AgentMiddlewareLayer:
|
||||
options: ChatOptions[ResponseModelBoundT],
|
||||
compaction_strategy: CompactionStrategy | None = None,
|
||||
tokenizer: TokenizerProtocol | None = None,
|
||||
function_invocation_kwargs: Mapping[str, Any] | None = None,
|
||||
client_kwargs: Mapping[str, Any] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Awaitable[AgentResponse[ResponseModelBoundT]]: ...
|
||||
|
||||
@@ -1129,6 +1167,8 @@ class AgentMiddlewareLayer:
|
||||
options: ChatOptions[None] | None = None,
|
||||
compaction_strategy: CompactionStrategy | None = None,
|
||||
tokenizer: TokenizerProtocol | None = None,
|
||||
function_invocation_kwargs: Mapping[str, Any] | None = None,
|
||||
client_kwargs: Mapping[str, Any] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Awaitable[AgentResponse[Any]]: ...
|
||||
|
||||
@@ -1143,6 +1183,8 @@ class AgentMiddlewareLayer:
|
||||
options: ChatOptions[Any] | None = None,
|
||||
compaction_strategy: CompactionStrategy | None = None,
|
||||
tokenizer: TokenizerProtocol | None = None,
|
||||
function_invocation_kwargs: Mapping[str, Any] | None = None,
|
||||
client_kwargs: Mapping[str, Any] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> ResponseStream[AgentResponseUpdate, AgentResponse[Any]]: ...
|
||||
|
||||
@@ -1156,6 +1198,8 @@ class AgentMiddlewareLayer:
|
||||
options: ChatOptions[Any] | None = None,
|
||||
compaction_strategy: CompactionStrategy | None = None,
|
||||
tokenizer: TokenizerProtocol | None = None,
|
||||
function_invocation_kwargs: Mapping[str, Any] | None = None,
|
||||
client_kwargs: Mapping[str, Any] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Awaitable[AgentResponse[Any]] | ResponseStream[AgentResponseUpdate, AgentResponse[Any]]:
|
||||
"""MiddlewareTypes-enabled unified run method."""
|
||||
@@ -1175,9 +1219,12 @@ class AgentMiddlewareLayer:
|
||||
+ run_middleware_list["function"]
|
||||
+ run_middleware_list["chat"]
|
||||
)
|
||||
combined_kwargs = dict(kwargs)
|
||||
combined_kwargs["middleware"] = combined_function_chat_middleware if combined_function_chat_middleware else None
|
||||
|
||||
effective_client_kwargs = dict(client_kwargs) if client_kwargs is not None else {}
|
||||
if combined_function_chat_middleware:
|
||||
effective_client_kwargs["middleware"] = combined_function_chat_middleware
|
||||
effective_function_invocation_kwargs = (
|
||||
dict(function_invocation_kwargs) if function_invocation_kwargs is not None else {}
|
||||
)
|
||||
# Execute with middleware if available
|
||||
if not pipeline.has_middlewares:
|
||||
return super().run( # type: ignore[misc, no-any-return]
|
||||
@@ -1187,7 +1234,9 @@ class AgentMiddlewareLayer:
|
||||
options=options,
|
||||
compaction_strategy=compaction_strategy,
|
||||
tokenizer=tokenizer,
|
||||
**combined_kwargs,
|
||||
function_invocation_kwargs=effective_function_invocation_kwargs,
|
||||
client_kwargs=effective_client_kwargs,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
context = AgentContext(
|
||||
@@ -1198,7 +1247,9 @@ class AgentMiddlewareLayer:
|
||||
stream=stream,
|
||||
compaction_strategy=compaction_strategy,
|
||||
tokenizer=tokenizer,
|
||||
kwargs=combined_kwargs,
|
||||
kwargs=kwargs,
|
||||
client_kwargs=effective_client_kwargs,
|
||||
function_invocation_kwargs=effective_function_invocation_kwargs,
|
||||
)
|
||||
|
||||
async def _execute() -> AgentResponse | ResponseStream[AgentResponseUpdate, AgentResponse] | None:
|
||||
@@ -1230,6 +1281,13 @@ class AgentMiddlewareLayer:
|
||||
def _middleware_handler(
|
||||
self, context: AgentContext
|
||||
) -> Awaitable[AgentResponse] | ResponseStream[AgentResponseUpdate, AgentResponse]:
|
||||
# TODO(Copilot): Delete once direct ``run(**kwargs)`` compatibility is removed.
|
||||
client_kwargs = {**context.client_kwargs, **context.kwargs}
|
||||
# TODO(Copilot): Delete once direct ``run(**kwargs)`` compatibility is removed.
|
||||
function_invocation_kwargs = {
|
||||
**context.function_invocation_kwargs,
|
||||
**{k: v for k, v in context.kwargs.items() if k != "middleware"},
|
||||
}
|
||||
return super().run( # type: ignore[misc, no-any-return]
|
||||
context.messages,
|
||||
stream=context.stream,
|
||||
@@ -1237,7 +1295,8 @@ class AgentMiddlewareLayer:
|
||||
options=context.options,
|
||||
compaction_strategy=context.compaction_strategy,
|
||||
tokenizer=context.tokenizer,
|
||||
**context.kwargs,
|
||||
function_invocation_kwargs=function_invocation_kwargs,
|
||||
client_kwargs=client_kwargs,
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -392,12 +392,16 @@ class BaseHistoryProvider(BaseContextProvider):
|
||||
self.store_outputs = store_outputs
|
||||
|
||||
@abstractmethod
|
||||
async def get_messages(self, session_id: str | None, **kwargs: Any) -> list[Message]:
|
||||
async def get_messages(
|
||||
self, session_id: str | None, *, state: dict[str, Any] | None = None, **kwargs: Any
|
||||
) -> list[Message]:
|
||||
"""Retrieve stored messages for this session.
|
||||
|
||||
Args:
|
||||
session_id: The session ID to retrieve messages for.
|
||||
**kwargs: Additional arguments (e.g., ``state`` for in-memory providers).
|
||||
state: Optional session state for providers that persist in session state.
|
||||
Not used by all providers.
|
||||
**kwargs: Additional subclass-specific extensibility arguments.
|
||||
|
||||
Returns:
|
||||
List of stored messages.
|
||||
@@ -405,13 +409,22 @@ class BaseHistoryProvider(BaseContextProvider):
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def save_messages(self, session_id: str | None, messages: Sequence[Message], **kwargs: Any) -> None:
|
||||
async def save_messages(
|
||||
self,
|
||||
session_id: str | None,
|
||||
messages: Sequence[Message],
|
||||
*,
|
||||
state: dict[str, Any] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
"""Persist messages for this session.
|
||||
|
||||
Args:
|
||||
session_id: The session ID to store messages for.
|
||||
messages: The messages to persist.
|
||||
**kwargs: Additional arguments (e.g., ``state`` for in-memory providers).
|
||||
state: Optional session state for providers that persist in session state.
|
||||
Not used by all providers.
|
||||
**kwargs: Additional subclass-specific extensibility arguments.
|
||||
"""
|
||||
...
|
||||
|
||||
|
||||
@@ -7,6 +7,8 @@ import inspect
|
||||
import json
|
||||
import logging
|
||||
import sys
|
||||
import typing
|
||||
import warnings
|
||||
from collections.abc import (
|
||||
AsyncIterable,
|
||||
Awaitable,
|
||||
@@ -37,7 +39,7 @@ from opentelemetry.metrics import Histogram, NoOpHistogram
|
||||
from pydantic import BaseModel, Field, ValidationError, create_model
|
||||
|
||||
from ._serialization import SerializationMixin
|
||||
from .exceptions import ToolException
|
||||
from .exceptions import ToolException, UserInputRequiredException
|
||||
from .observability import (
|
||||
OPERATION_DURATION_BUCKET_BOUNDARIES,
|
||||
OtelAttr,
|
||||
@@ -61,7 +63,8 @@ if TYPE_CHECKING:
|
||||
from ._clients import SupportsChatGetResponse
|
||||
from ._compaction import CompactionStrategy, TokenizerProtocol
|
||||
from ._mcp import MCPTool
|
||||
from ._middleware import FunctionMiddlewarePipeline, FunctionMiddlewareTypes
|
||||
from ._middleware import FunctionInvocationContext, FunctionMiddlewarePipeline, FunctionMiddlewareTypes
|
||||
from ._sessions import AgentSession
|
||||
from ._types import (
|
||||
ChatOptions,
|
||||
ChatResponse,
|
||||
@@ -71,7 +74,6 @@ if TYPE_CHECKING:
|
||||
ResponseStream,
|
||||
)
|
||||
|
||||
ResponseModelBoundT = TypeVar("ResponseModelBoundT", bound=BaseModel)
|
||||
else:
|
||||
MCPTool = Any # type: ignore[assignment,misc]
|
||||
|
||||
@@ -83,9 +85,23 @@ DEFAULT_MAX_ITERATIONS: Final[int] = 40
|
||||
DEFAULT_MAX_CONSECUTIVE_ERRORS_PER_REQUEST: Final[int] = 3
|
||||
SHELL_TOOL_KIND_VALUE: Final[str] = "shell"
|
||||
ChatClientT = TypeVar("ChatClientT", bound="SupportsChatGetResponse[Any]")
|
||||
ResponseModelBoundT = TypeVar("ResponseModelBoundT", bound=BaseModel)
|
||||
|
||||
# region Helpers
|
||||
|
||||
|
||||
def _get_tool_name(tool: Any) -> str | None:
|
||||
"""Extract a tool name from a tool object or dict tool definition."""
|
||||
if isinstance(tool, Mapping):
|
||||
func = tool.get("function", None) # type: ignore
|
||||
if func and isinstance(func, Mapping):
|
||||
name = func.get("name") # type: ignore
|
||||
return name if isinstance(name, str) else None
|
||||
return None
|
||||
name = getattr(tool, "name", None)
|
||||
return name if isinstance(name, str) else None
|
||||
|
||||
|
||||
def _parse_inputs( # pyright: ignore[reportUnusedFunction]
|
||||
inputs: Content | dict[str, Any] | str | list[Content | dict[str, Any] | str] | None,
|
||||
) -> list[Content]:
|
||||
@@ -174,6 +190,16 @@ def _default_histogram() -> Histogram:
|
||||
)
|
||||
|
||||
|
||||
def _annotation_includes_function_invocation_context(annotation: Any) -> bool:
|
||||
"""Check whether an annotation resolves to FunctionInvocationContext."""
|
||||
from ._middleware import FunctionInvocationContext
|
||||
|
||||
candidates = get_args(annotation) or (annotation,)
|
||||
return any(
|
||||
candidate is FunctionInvocationContext or candidate == "FunctionInvocationContext" for candidate in candidates
|
||||
)
|
||||
|
||||
|
||||
ClassT = TypeVar("ClassT", bound="SerializationMixin")
|
||||
|
||||
|
||||
@@ -310,6 +336,12 @@ class FunctionTool(SerializationMixin):
|
||||
# FunctionTool-specific attributes
|
||||
self.func = func
|
||||
self._instance = None # Store the instance for bound methods
|
||||
self._context_parameter_name: str | None = None
|
||||
self._input_model_explicitly_provided = input_model is not None
|
||||
# TODO(Copilot): Delete once legacy ``**kwargs`` runtime injection is removed.
|
||||
self._forward_runtime_kwargs: bool = False
|
||||
if self.func:
|
||||
self._discover_injected_parameters()
|
||||
|
||||
# Initialize schema cache (will be lazily populated)
|
||||
self._input_schema_cached: dict[str, Any] | None = None
|
||||
@@ -336,13 +368,37 @@ class FunctionTool(SerializationMixin):
|
||||
self._invocation_duration_histogram = _default_histogram()
|
||||
self.type: Literal["function_tool"] = "function_tool"
|
||||
self.result_parser = result_parser
|
||||
self._forward_runtime_kwargs: bool = False
|
||||
if self.func:
|
||||
sig = inspect.signature(self.func)
|
||||
for param in sig.parameters.values():
|
||||
if param.kind == inspect.Parameter.VAR_KEYWORD:
|
||||
self._forward_runtime_kwargs = True
|
||||
break
|
||||
|
||||
def _discover_injected_parameters(self) -> None:
|
||||
"""Inspect the wrapped function for runtime injection parameters."""
|
||||
func = self.func.func if isinstance(self.func, FunctionTool) else self.func
|
||||
if func is None:
|
||||
return
|
||||
|
||||
signature = inspect.signature(func)
|
||||
try:
|
||||
type_hints = typing.get_type_hints(func)
|
||||
except Exception:
|
||||
type_hints = {name: param.annotation for name, param in signature.parameters.items()}
|
||||
|
||||
for name, param in signature.parameters.items():
|
||||
if name in {"self", "cls"}:
|
||||
continue
|
||||
if param.kind == inspect.Parameter.VAR_KEYWORD:
|
||||
self._forward_runtime_kwargs = True
|
||||
continue
|
||||
|
||||
annotation = type_hints.get(name, param.annotation)
|
||||
if self._is_context_parameter(name, annotation):
|
||||
if self._context_parameter_name is not None:
|
||||
raise ValueError(f"Function '{self.name}' defines multiple FunctionInvocationContext parameters.")
|
||||
self._context_parameter_name = name
|
||||
|
||||
def _is_context_parameter(self, name: str, annotation: Any) -> bool:
|
||||
"""Check whether a callable parameter should receive FunctionInvocationContext injection."""
|
||||
if _annotation_includes_function_invocation_context(annotation):
|
||||
return True
|
||||
return self._input_model_explicitly_provided and name == "ctx" and annotation is inspect.Parameter.empty
|
||||
|
||||
def __str__(self) -> str:
|
||||
"""Return a string representation of the tool."""
|
||||
@@ -411,6 +467,7 @@ class FunctionTool(SerializationMixin):
|
||||
)
|
||||
for pname, param in sig.parameters.items()
|
||||
if pname not in {"self", "cls"}
|
||||
and pname != self._context_parameter_name
|
||||
and param.kind not in {inspect.Parameter.VAR_POSITIONAL, inspect.Parameter.VAR_KEYWORD}
|
||||
}
|
||||
return create_model(f"{self.name}_input", **fields)
|
||||
@@ -448,6 +505,7 @@ class FunctionTool(SerializationMixin):
|
||||
self,
|
||||
*,
|
||||
arguments: BaseModel | Mapping[str, Any] | None = None,
|
||||
context: FunctionInvocationContext | None = None,
|
||||
**kwargs: Any,
|
||||
) -> list[Content]:
|
||||
"""Run the AI function with the provided arguments as a Pydantic model.
|
||||
@@ -459,7 +517,8 @@ class FunctionTool(SerializationMixin):
|
||||
|
||||
Keyword Args:
|
||||
arguments: A mapping or model instance containing the arguments for the function.
|
||||
kwargs: Keyword arguments to pass to the function, will not be used if ``arguments`` is provided.
|
||||
context: Explicit function invocation context carrying runtime kwargs.
|
||||
kwargs: Deprecated keyword arguments to pass to the function. Use ``context`` instead.
|
||||
|
||||
Returns:
|
||||
A list of Content items representing the tool output.
|
||||
@@ -470,14 +529,37 @@ class FunctionTool(SerializationMixin):
|
||||
if self.declaration_only:
|
||||
raise ToolException(f"Function '{self.name}' is declaration only and cannot be invoked.")
|
||||
global OBSERVABILITY_SETTINGS
|
||||
from ._middleware import FunctionInvocationContext
|
||||
from ._types import Content
|
||||
from .observability import OBSERVABILITY_SETTINGS
|
||||
|
||||
parser = self.result_parser or FunctionTool.parse_result
|
||||
|
||||
original_kwargs = dict(kwargs)
|
||||
tool_call_id = original_kwargs.pop("tool_call_id", None)
|
||||
if arguments is not None:
|
||||
parameter_names = set(self.parameters().get("properties", {}).keys())
|
||||
direct_argument_kwargs = (
|
||||
{key: value for key, value in kwargs.items() if key in parameter_names} if arguments is None else {}
|
||||
)
|
||||
runtime_kwargs = dict(context.kwargs) if context is not None else {}
|
||||
deprecated_runtime_kwargs = {
|
||||
key: value for key, value in kwargs.items() if key not in direct_argument_kwargs and key != "tool_call_id"
|
||||
}
|
||||
if deprecated_runtime_kwargs:
|
||||
warnings.warn(
|
||||
"Passing runtime keyword arguments directly to FunctionTool.invoke() is deprecated; "
|
||||
"pass them via FunctionInvocationContext instead.",
|
||||
DeprecationWarning,
|
||||
stacklevel=2,
|
||||
)
|
||||
runtime_kwargs.update(deprecated_runtime_kwargs)
|
||||
tool_call_id = kwargs.get("tool_call_id", runtime_kwargs.pop("tool_call_id", None))
|
||||
if arguments is None and direct_argument_kwargs:
|
||||
arguments = direct_argument_kwargs
|
||||
if arguments is None and context is not None:
|
||||
arguments = context.arguments
|
||||
|
||||
if arguments is None:
|
||||
validated_arguments: dict[str, Any] = {}
|
||||
else:
|
||||
try:
|
||||
if isinstance(arguments, Mapping):
|
||||
parsed_arguments = dict(arguments)
|
||||
@@ -499,19 +581,45 @@ class FunctionTool(SerializationMixin):
|
||||
)
|
||||
except ValidationError as exc:
|
||||
raise TypeError(f"Invalid arguments for '{self.name}': {exc}") from exc
|
||||
kwargs = _validate_arguments_against_schema(
|
||||
|
||||
validated_arguments = _validate_arguments_against_schema(
|
||||
arguments=parsed_arguments,
|
||||
schema=self.parameters(),
|
||||
tool_name=self.name,
|
||||
)
|
||||
if getattr(self, "_forward_runtime_kwargs", False) and original_kwargs:
|
||||
kwargs.update(original_kwargs)
|
||||
else:
|
||||
kwargs = original_kwargs
|
||||
|
||||
effective_context = context
|
||||
if effective_context is None and self._context_parameter_name is not None:
|
||||
effective_context = FunctionInvocationContext(
|
||||
function=self,
|
||||
arguments=validated_arguments,
|
||||
kwargs=runtime_kwargs,
|
||||
)
|
||||
if effective_context is not None:
|
||||
effective_context.function = self
|
||||
effective_context.arguments = validated_arguments
|
||||
effective_context.kwargs = dict(runtime_kwargs)
|
||||
|
||||
call_kwargs = dict(validated_arguments)
|
||||
observable_kwargs = dict(validated_arguments)
|
||||
|
||||
# Legacy runtime kwargs injection path retained for backwards compatibility with tools
|
||||
# that still declare ``**kwargs``. New tools should consume runtime data via ``ctx``.
|
||||
legacy_runtime_kwargs = dict(runtime_kwargs)
|
||||
if self._forward_runtime_kwargs and legacy_runtime_kwargs:
|
||||
for key, value in legacy_runtime_kwargs.items():
|
||||
if key not in call_kwargs:
|
||||
call_kwargs[key] = value
|
||||
if key not in observable_kwargs:
|
||||
observable_kwargs[key] = value
|
||||
|
||||
if self._context_parameter_name is not None and effective_context is not None:
|
||||
call_kwargs[self._context_parameter_name] = effective_context
|
||||
|
||||
if not OBSERVABILITY_SETTINGS.ENABLED: # type: ignore[name-defined]
|
||||
logger.info(f"Function name: {self.name}")
|
||||
logger.debug(f"Function arguments: {kwargs}")
|
||||
res = self.__call__(**kwargs)
|
||||
logger.debug(f"Function arguments: {observable_kwargs}")
|
||||
res = self.__call__(**call_kwargs)
|
||||
result = await res if inspect.isawaitable(res) else res
|
||||
try:
|
||||
parsed = parser(result)
|
||||
@@ -532,7 +640,7 @@ class FunctionTool(SerializationMixin):
|
||||
# Filter out framework kwargs that are not JSON serializable.
|
||||
serializable_kwargs = {
|
||||
k: v
|
||||
for k, v in kwargs.items()
|
||||
for k, v in observable_kwargs.items()
|
||||
if k
|
||||
not in {
|
||||
"chat_options",
|
||||
@@ -558,7 +666,7 @@ class FunctionTool(SerializationMixin):
|
||||
start_time_stamp = perf_counter()
|
||||
end_time_stamp: float | None = None
|
||||
try:
|
||||
res = self.__call__(**kwargs)
|
||||
res = self.__call__(**call_kwargs)
|
||||
result = await res if inspect.isawaitable(res) else res
|
||||
end_time_stamp = perf_counter()
|
||||
except Exception as exception:
|
||||
@@ -701,6 +809,51 @@ class FunctionTool(SerializationMixin):
|
||||
ToolTypes: TypeAlias = FunctionTool | MCPTool | Mapping[str, Any] | object
|
||||
|
||||
|
||||
def _raise_duplicate_tool_name(tool_name: str, duplicate_error_message: str | None = None) -> None:
|
||||
message = duplicate_error_message or "Tool names must be unique."
|
||||
raise ValueError(f"Duplicate tool name '{tool_name}'. {message}")
|
||||
|
||||
|
||||
def _append_unique_tools(
|
||||
existing_tools: list[ToolTypes],
|
||||
new_tools: Sequence[ToolTypes],
|
||||
*,
|
||||
duplicate_error_message: str | None = None,
|
||||
) -> list[ToolTypes]:
|
||||
seen_by_name: dict[str, ToolTypes] = {}
|
||||
for tool_item in existing_tools:
|
||||
if tool_name := _get_tool_name(tool_item):
|
||||
seen_by_name[tool_name] = tool_item
|
||||
|
||||
for tool_item in new_tools:
|
||||
tool_name = _get_tool_name(tool_item)
|
||||
if tool_name is None:
|
||||
existing_tools.append(tool_item)
|
||||
continue
|
||||
|
||||
existing_tool = seen_by_name.get(tool_name)
|
||||
if existing_tool is None:
|
||||
seen_by_name[tool_name] = tool_item
|
||||
existing_tools.append(tool_item)
|
||||
continue
|
||||
|
||||
if existing_tool is tool_item:
|
||||
continue
|
||||
|
||||
_raise_duplicate_tool_name(tool_name, duplicate_error_message)
|
||||
|
||||
return existing_tools
|
||||
|
||||
|
||||
def _ensure_unique_tool_names(
|
||||
tools: ToolTypes | Callable[..., Any] | Sequence[ToolTypes | Callable[..., Any]],
|
||||
*,
|
||||
duplicate_error_message: str | None = None,
|
||||
) -> list[ToolTypes]:
|
||||
normalized_tools = normalize_tools(tools)
|
||||
return _append_unique_tools([], normalized_tools, duplicate_error_message=duplicate_error_message)
|
||||
|
||||
|
||||
def normalize_tools(
|
||||
tools: ToolTypes | Callable[..., Any] | Sequence[ToolTypes | Callable[..., Any]] | None,
|
||||
) -> list[ToolTypes]:
|
||||
@@ -1160,9 +1313,10 @@ async def _auto_invoke_function(
|
||||
*,
|
||||
config: FunctionInvocationConfiguration,
|
||||
tool_map: dict[str, FunctionTool],
|
||||
invocation_session: AgentSession | None = None,
|
||||
sequence_index: int | None = None,
|
||||
request_index: int | None = None,
|
||||
middleware_pipeline: FunctionMiddlewarePipeline | None = None, # Optional MiddlewarePipeline
|
||||
middleware_pipeline: FunctionMiddlewarePipeline | None = None,
|
||||
) -> Content:
|
||||
"""Invoke a function call requested by the agent, applying middleware that is defined.
|
||||
|
||||
@@ -1173,6 +1327,7 @@ async def _auto_invoke_function(
|
||||
Keyword Args:
|
||||
config: The function invocation configuration.
|
||||
tool_map: A mapping of tool names to FunctionTool instances.
|
||||
invocation_session: The agent session for this invocation, if any.
|
||||
sequence_index: The index of the function call in the sequence.
|
||||
request_index: The index of the request iteration.
|
||||
middleware_pipeline: Optional middleware pipeline to apply during execution.
|
||||
@@ -1224,6 +1379,8 @@ async def _auto_invoke_function(
|
||||
for key, value in (custom_args or {}).items()
|
||||
if key not in {"_function_middleware_pipeline", "middleware", "conversation_id"}
|
||||
}
|
||||
if invocation_session is not None:
|
||||
runtime_kwargs["session"] = invocation_session
|
||||
try:
|
||||
if not cast(bool, getattr(tool, "_schema_supplied", False)) and tool.input_model is not None:
|
||||
args = tool.input_model.model_validate(parsed_args).model_dump(exclude_none=True)
|
||||
@@ -1245,19 +1402,31 @@ async def _auto_invoke_function(
|
||||
additional_properties=function_call_content.additional_properties,
|
||||
)
|
||||
|
||||
from ._middleware import FunctionInvocationContext
|
||||
|
||||
if middleware_pipeline is None or not middleware_pipeline.has_middlewares:
|
||||
# No middleware - execute directly
|
||||
try:
|
||||
direct_context = None
|
||||
if getattr(tool, "_forward_runtime_kwargs", False) or getattr(tool, "_context_parameter_name", None):
|
||||
direct_context = FunctionInvocationContext(
|
||||
function=tool,
|
||||
arguments=args,
|
||||
session=invocation_session,
|
||||
kwargs=runtime_kwargs.copy(),
|
||||
)
|
||||
function_result = await tool.invoke(
|
||||
arguments=args,
|
||||
context=direct_context,
|
||||
tool_call_id=function_call_content.call_id,
|
||||
**runtime_kwargs if getattr(tool, "_forward_runtime_kwargs", False) else {},
|
||||
)
|
||||
return Content.from_function_result(
|
||||
call_id=function_call_content.call_id, # type: ignore[arg-type]
|
||||
result=function_result,
|
||||
additional_properties=function_call_content.additional_properties,
|
||||
)
|
||||
except UserInputRequiredException:
|
||||
raise
|
||||
except Exception as exc:
|
||||
message = "Error: Function failed."
|
||||
if config.get("include_detailed_errors", False):
|
||||
@@ -1269,19 +1438,18 @@ async def _auto_invoke_function(
|
||||
additional_properties=function_call_content.additional_properties,
|
||||
)
|
||||
# Execute through middleware pipeline if available
|
||||
from ._middleware import FunctionInvocationContext
|
||||
|
||||
middleware_context = FunctionInvocationContext(
|
||||
function=tool,
|
||||
arguments=args,
|
||||
session=invocation_session,
|
||||
kwargs=runtime_kwargs.copy(),
|
||||
)
|
||||
|
||||
async def final_function_handler(context_obj: Any) -> Any:
|
||||
return await tool.invoke(
|
||||
arguments=context_obj.arguments,
|
||||
context=context_obj,
|
||||
tool_call_id=function_call_content.call_id,
|
||||
**context_obj.kwargs if getattr(tool, "_forward_runtime_kwargs", False) else {},
|
||||
)
|
||||
|
||||
from ._middleware import MiddlewareTermination
|
||||
@@ -1304,6 +1472,8 @@ async def _auto_invoke_function(
|
||||
additional_properties=function_call_content.additional_properties,
|
||||
)
|
||||
raise
|
||||
except UserInputRequiredException:
|
||||
raise
|
||||
except Exception as exc:
|
||||
message = "Error: Function failed."
|
||||
if config.get("include_detailed_errors", False):
|
||||
@@ -1320,7 +1490,7 @@ def _get_tool_map(
|
||||
tools: ToolTypes | Callable[..., Any] | Sequence[ToolTypes | Callable[..., Any]],
|
||||
) -> dict[str, FunctionTool]:
|
||||
tool_list: dict[str, FunctionTool] = {}
|
||||
for tool_item in normalize_tools(tools):
|
||||
for tool_item in _ensure_unique_tool_names(tools):
|
||||
if isinstance(tool_item, FunctionTool):
|
||||
tool_list[tool_item.name] = tool_item
|
||||
return tool_list
|
||||
@@ -1332,7 +1502,8 @@ async def _try_execute_function_calls(
|
||||
function_calls: Sequence[Content],
|
||||
tools: ToolTypes | Callable[..., Any] | Sequence[ToolTypes | Callable[..., Any]],
|
||||
config: FunctionInvocationConfiguration,
|
||||
middleware_pipeline: Any = None, # Optional MiddlewarePipeline to avoid circular imports
|
||||
invocation_session: AgentSession | None = None,
|
||||
middleware_pipeline: Any = None,
|
||||
) -> tuple[Sequence[Content], bool]:
|
||||
"""Execute multiple function calls concurrently.
|
||||
|
||||
@@ -1342,6 +1513,7 @@ async def _try_execute_function_calls(
|
||||
function_calls: A sequence of FunctionCallContent to execute.
|
||||
tools: The tools available for execution.
|
||||
config: Configuration for function invocation.
|
||||
invocation_session: The agent session for this invocation, if any.
|
||||
middleware_pipeline: Optional middleware pipeline to apply during execution.
|
||||
|
||||
Returns:
|
||||
@@ -1411,6 +1583,8 @@ async def _try_execute_function_calls(
|
||||
# Run all function calls concurrently, handling MiddlewareTermination
|
||||
from ._middleware import MiddlewareTermination
|
||||
|
||||
extra_user_input_contents: list[Content] = []
|
||||
|
||||
async def invoke_with_termination_handling(
|
||||
function_call: Content,
|
||||
seq_idx: int,
|
||||
@@ -1421,6 +1595,7 @@ async def _try_execute_function_calls(
|
||||
function_call_content=function_call, # type: ignore[arg-type]
|
||||
custom_args=custom_args,
|
||||
tool_map=tool_map,
|
||||
invocation_session=invocation_session,
|
||||
sequence_index=seq_idx,
|
||||
request_index=attempt_idx,
|
||||
middleware_pipeline=middleware_pipeline,
|
||||
@@ -1437,6 +1612,26 @@ async def _try_execute_function_calls(
|
||||
result=exc.result,
|
||||
)
|
||||
return (result_content, True)
|
||||
except UserInputRequiredException as exc:
|
||||
if exc.contents:
|
||||
propagated: list[Content] = []
|
||||
for item in exc.contents:
|
||||
if isinstance(item, Content):
|
||||
item.call_id = function_call.call_id # type: ignore[attr-defined]
|
||||
if not item.id: # type: ignore[attr-defined]
|
||||
item.id = function_call.call_id # type: ignore[attr-defined]
|
||||
propagated.append(item)
|
||||
if propagated:
|
||||
extra_user_input_contents.extend(propagated[1:])
|
||||
return (propagated[0], False)
|
||||
return (
|
||||
Content.from_function_result(
|
||||
call_id=function_call.call_id, # type: ignore[arg-type]
|
||||
result="Tool requires user input but no request details were provided.",
|
||||
exception="UserInputRequiredException",
|
||||
),
|
||||
False,
|
||||
)
|
||||
|
||||
execution_results = await asyncio.gather(*[
|
||||
invoke_with_termination_handling(function_call, seq_idx) for seq_idx, function_call in enumerate(function_calls)
|
||||
@@ -1444,6 +1639,7 @@ async def _try_execute_function_calls(
|
||||
|
||||
# Unpack results - each is (Content, terminate_flag)
|
||||
contents: list[Content] = [result[0] for result in execution_results]
|
||||
contents.extend(extra_user_input_contents)
|
||||
# If any function requested termination, terminate the loop
|
||||
should_terminate = any(result[1] for result in execution_results)
|
||||
return (contents, should_terminate)
|
||||
@@ -1456,6 +1652,7 @@ async def _execute_function_calls(
|
||||
function_calls: list[Content],
|
||||
tool_options: dict[str, Any] | None,
|
||||
config: FunctionInvocationConfiguration,
|
||||
invocation_session: AgentSession | None = None,
|
||||
middleware_pipeline: Any = None,
|
||||
) -> tuple[list[Content], bool, bool]:
|
||||
tools = _extract_tools(tool_options)
|
||||
@@ -1466,6 +1663,7 @@ async def _execute_function_calls(
|
||||
attempt_idx=attempt_idx,
|
||||
function_calls=function_calls,
|
||||
tools=tools, # type: ignore
|
||||
invocation_session=invocation_session,
|
||||
middleware_pipeline=middleware_pipeline,
|
||||
config=config,
|
||||
)
|
||||
@@ -1675,7 +1873,10 @@ def _handle_function_call_results(
|
||||
) -> FunctionRequestResult:
|
||||
from ._types import Message
|
||||
|
||||
if any(fccr.type in {"function_approval_request", "function_call"} for fccr in function_call_results):
|
||||
if any(
|
||||
fccr.type in {"function_approval_request", "function_call"} or fccr.user_input_request
|
||||
for fccr in function_call_results
|
||||
):
|
||||
# Only add items that aren't already in the message (e.g. function_approval_request wrappers).
|
||||
# Declaration-only function_call items are already present from the LLM response.
|
||||
new_items = [fccr for fccr in function_call_results if fccr.type != "function_call"]
|
||||
@@ -1843,6 +2044,8 @@ class FunctionInvocationLayer(Generic[OptionsCoT]):
|
||||
options: ChatOptions[ResponseModelBoundT],
|
||||
compaction_strategy: CompactionStrategy | None = None,
|
||||
tokenizer: TokenizerProtocol | None = None,
|
||||
function_invocation_kwargs: Mapping[str, Any] | None = None,
|
||||
client_kwargs: Mapping[str, Any] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Awaitable[ChatResponse[ResponseModelBoundT]]: ...
|
||||
|
||||
@@ -1855,6 +2058,8 @@ class FunctionInvocationLayer(Generic[OptionsCoT]):
|
||||
options: OptionsCoT | ChatOptions[None] | None = None,
|
||||
compaction_strategy: CompactionStrategy | None = None,
|
||||
tokenizer: TokenizerProtocol | None = None,
|
||||
function_invocation_kwargs: Mapping[str, Any] | None = None,
|
||||
client_kwargs: Mapping[str, Any] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Awaitable[ChatResponse[Any]]: ...
|
||||
|
||||
@@ -1867,6 +2072,8 @@ class FunctionInvocationLayer(Generic[OptionsCoT]):
|
||||
options: OptionsCoT | ChatOptions[Any] | None = None,
|
||||
compaction_strategy: CompactionStrategy | None = None,
|
||||
tokenizer: TokenizerProtocol | None = None,
|
||||
function_invocation_kwargs: Mapping[str, Any] | None = None,
|
||||
client_kwargs: Mapping[str, Any] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> ResponseStream[ChatResponseUpdate, ChatResponse[Any]]: ...
|
||||
|
||||
@@ -1879,6 +2086,8 @@ class FunctionInvocationLayer(Generic[OptionsCoT]):
|
||||
function_middleware: Sequence[FunctionMiddlewareTypes] | None = None,
|
||||
compaction_strategy: CompactionStrategy | None = None,
|
||||
tokenizer: TokenizerProtocol | None = None,
|
||||
function_invocation_kwargs: Mapping[str, Any] | None = None,
|
||||
client_kwargs: Mapping[str, Any] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Awaitable[ChatResponse[Any]] | ResponseStream[ChatResponseUpdate, ChatResponse[Any]]:
|
||||
from ._middleware import FunctionMiddlewarePipeline
|
||||
@@ -1889,28 +2098,45 @@ class FunctionInvocationLayer(Generic[OptionsCoT]):
|
||||
)
|
||||
|
||||
super_get_response = super().get_response # type: ignore[misc]
|
||||
if kwargs:
|
||||
warnings.warn(
|
||||
"Passing client-specific keyword arguments directly to get_response() is deprecated; "
|
||||
"pass them via client_kwargs instead.",
|
||||
DeprecationWarning,
|
||||
stacklevel=2,
|
||||
)
|
||||
|
||||
effective_client_kwargs = dict(client_kwargs) if client_kwargs is not None else {}
|
||||
effective_function_middleware = function_middleware
|
||||
if effective_function_middleware is None:
|
||||
middleware_from_client_kwargs = effective_client_kwargs.pop("function_middleware", None)
|
||||
if middleware_from_client_kwargs is not None:
|
||||
effective_function_middleware = cast(Sequence[Any], middleware_from_client_kwargs)
|
||||
|
||||
# ChatMiddleware adds this kwarg
|
||||
function_middleware_pipeline = FunctionMiddlewarePipeline(
|
||||
*(self.function_middleware), *(function_middleware or [])
|
||||
*(self.function_middleware), *(effective_function_middleware or [])
|
||||
)
|
||||
max_errors = self.function_invocation_configuration.get(
|
||||
"max_consecutive_errors_per_request", DEFAULT_MAX_CONSECUTIVE_ERRORS_PER_REQUEST
|
||||
)
|
||||
additional_function_arguments: dict[str, Any] = {}
|
||||
additional_function_arguments = (
|
||||
dict(function_invocation_kwargs) if function_invocation_kwargs is not None else {}
|
||||
)
|
||||
if options and (additional_opts := options.get("additional_function_arguments")): # type: ignore[attr-defined]
|
||||
additional_function_arguments = additional_opts # type: ignore
|
||||
additional_function_arguments.update(cast(Mapping[str, Any], additional_opts))
|
||||
from ._sessions import AgentSession as _AgentSession
|
||||
|
||||
raw_session = effective_client_kwargs.get("session")
|
||||
invocation_session = raw_session if isinstance(raw_session, _AgentSession) else None
|
||||
execute_function_calls = partial(
|
||||
_execute_function_calls,
|
||||
custom_args=additional_function_arguments,
|
||||
config=self.function_invocation_configuration,
|
||||
invocation_session=invocation_session,
|
||||
middleware_pipeline=function_middleware_pipeline,
|
||||
)
|
||||
filtered_kwargs = {k: v for k, v in kwargs.items() if k != "session"}
|
||||
if compaction_strategy is not None:
|
||||
filtered_kwargs["compaction_strategy"] = compaction_strategy
|
||||
if tokenizer is not None:
|
||||
filtered_kwargs["tokenizer"] = tokenizer
|
||||
filtered_kwargs = {k: v for k, v in {**effective_client_kwargs, **kwargs}.items() if k != "session"}
|
||||
|
||||
# Make options mutable so we can update conversation_id during function invocation loop
|
||||
mutable_options: dict[str, Any] = dict(options) if options else {}
|
||||
@@ -1960,7 +2186,9 @@ class FunctionInvocationLayer(Generic[OptionsCoT]):
|
||||
messages=prepped_messages,
|
||||
stream=False,
|
||||
options=mutable_options,
|
||||
**filtered_kwargs,
|
||||
compaction_strategy=compaction_strategy,
|
||||
tokenizer=tokenizer,
|
||||
client_kwargs=filtered_kwargs,
|
||||
),
|
||||
)
|
||||
|
||||
@@ -2029,7 +2257,9 @@ class FunctionInvocationLayer(Generic[OptionsCoT]):
|
||||
messages=prepped_messages,
|
||||
stream=False,
|
||||
options=mutable_options,
|
||||
**filtered_kwargs,
|
||||
compaction_strategy=compaction_strategy,
|
||||
tokenizer=tokenizer,
|
||||
client_kwargs=filtered_kwargs,
|
||||
),
|
||||
)
|
||||
if fcc_messages:
|
||||
@@ -2079,7 +2309,9 @@ class FunctionInvocationLayer(Generic[OptionsCoT]):
|
||||
messages=prepped_messages,
|
||||
stream=True,
|
||||
options=mutable_options,
|
||||
**filtered_kwargs,
|
||||
compaction_strategy=compaction_strategy,
|
||||
tokenizer=tokenizer,
|
||||
client_kwargs=filtered_kwargs,
|
||||
),
|
||||
)
|
||||
await inner_stream
|
||||
@@ -2171,7 +2403,9 @@ class FunctionInvocationLayer(Generic[OptionsCoT]):
|
||||
messages=prepped_messages,
|
||||
stream=True,
|
||||
options=mutable_options,
|
||||
**filtered_kwargs,
|
||||
compaction_strategy=compaction_strategy,
|
||||
tokenizer=tokenizer,
|
||||
client_kwargs=filtered_kwargs,
|
||||
),
|
||||
)
|
||||
await final_inner_stream
|
||||
|
||||
@@ -2698,7 +2698,7 @@ class ResponseStream(AsyncIterable[UpdateT], Generic[UpdateT, FinalT]):
|
||||
stream: AsyncIterable[UpdateT] | Awaitable[AsyncIterable[UpdateT]],
|
||||
*,
|
||||
finalizer: Callable[[Sequence[UpdateT]], FinalT | Awaitable[FinalT]] | None = None,
|
||||
transform_hooks: list[Callable[[UpdateT], UpdateT | Awaitable[UpdateT] | None]] | None = None,
|
||||
transform_hooks: list[Callable[[UpdateT], UpdateT | Awaitable[UpdateT | None] | None]] | None = None,
|
||||
cleanup_hooks: list[Callable[[], Awaitable[None] | None]] | None = None,
|
||||
result_hooks: list[Callable[[FinalT], FinalT | Awaitable[FinalT | None] | None]] | None = None,
|
||||
) -> None:
|
||||
@@ -2722,7 +2722,7 @@ class ResponseStream(AsyncIterable[UpdateT], Generic[UpdateT, FinalT]):
|
||||
self._consumed: bool = False
|
||||
self._finalized: bool = False
|
||||
self._final_result: FinalT | None = None
|
||||
self._transform_hooks: list[Callable[[UpdateT], UpdateT | Awaitable[UpdateT] | None]] = (
|
||||
self._transform_hooks: list[Callable[[UpdateT], UpdateT | Awaitable[UpdateT | None] | None]] = (
|
||||
transform_hooks if transform_hooks is not None else []
|
||||
)
|
||||
self._result_hooks: list[Callable[[FinalT], FinalT | Awaitable[FinalT | None] | None]] = (
|
||||
@@ -2995,7 +2995,7 @@ class ResponseStream(AsyncIterable[UpdateT], Generic[UpdateT, FinalT]):
|
||||
|
||||
def with_transform_hook(
|
||||
self,
|
||||
hook: Callable[[UpdateT], UpdateT | Awaitable[UpdateT] | None],
|
||||
hook: Callable[[UpdateT], UpdateT | Awaitable[UpdateT | None] | None],
|
||||
) -> ResponseStream[UpdateT, FinalT]:
|
||||
"""Register a transform hook executed for each update during iteration."""
|
||||
self._transform_hooks.append(hook)
|
||||
|
||||
@@ -9,6 +9,7 @@ from collections.abc import Awaitable, Callable, Sequence
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, ClassVar, TypeAlias, TypeVar
|
||||
|
||||
from .._agents import SupportsAgentRun
|
||||
from ._const import INTERNAL_SOURCE_ID
|
||||
from ._executor import Executor
|
||||
from ._model_utils import DictConvertible, encode_value
|
||||
@@ -264,7 +265,7 @@ class Case:
|
||||
"""
|
||||
|
||||
condition: Callable[[Any], bool]
|
||||
target: Executor | str
|
||||
target: Executor | SupportsAgentRun
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -287,7 +288,7 @@ class Default:
|
||||
assert fallback.target.id == "dead_letter"
|
||||
"""
|
||||
|
||||
target: Executor | str
|
||||
target: Executor | SupportsAgentRun
|
||||
|
||||
|
||||
@dataclass(init=False)
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
# ruff: noqa: RUF070, RUF100
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
|
||||
@@ -172,12 +172,12 @@ class AzureOpenAIChatClient( # type: ignore[misc]
|
||||
credential: AzureCredentialTypes | AzureTokenProvider | None = None,
|
||||
default_headers: Mapping[str, str] | None = None,
|
||||
async_client: AsyncAzureOpenAI | None = None,
|
||||
additional_properties: dict[str, Any] | None = None,
|
||||
env_file_path: str | None = None,
|
||||
env_file_encoding: str | None = None,
|
||||
instruction_role: str | None = None,
|
||||
middleware: Sequence[MiddlewareTypes] | None = None,
|
||||
function_invocation_configuration: FunctionInvocationConfiguration | None = None,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
"""Initialize an Azure OpenAI Chat completion client.
|
||||
|
||||
@@ -205,13 +205,13 @@ class AzureOpenAIChatClient( # type: ignore[misc]
|
||||
default_headers: The default headers mapping of string keys to
|
||||
string values for HTTP requests.
|
||||
async_client: An existing client to use.
|
||||
additional_properties: Additional properties stored on the client instance.
|
||||
env_file_path: Use the environment settings file as a fallback to using env vars.
|
||||
env_file_encoding: The encoding of the environment settings file, defaults to 'utf-8'.
|
||||
instruction_role: The role to use for 'instruction' messages, for example, summarization
|
||||
prompts could use `developer` or `system`.
|
||||
middleware: Optional sequence of middleware to apply to requests.
|
||||
function_invocation_configuration: Optional configuration for function invocation behavior.
|
||||
kwargs: Other keyword parameters.
|
||||
|
||||
Examples:
|
||||
.. code-block:: python
|
||||
@@ -283,10 +283,10 @@ class AzureOpenAIChatClient( # type: ignore[misc]
|
||||
credential=credential,
|
||||
default_headers=default_headers,
|
||||
client=async_client,
|
||||
additional_properties=additional_properties,
|
||||
instruction_role=instruction_role,
|
||||
middleware=middleware,
|
||||
function_invocation_configuration=function_invocation_configuration,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
@override
|
||||
|
||||
@@ -180,6 +180,34 @@ class ToolExecutionException(ToolException):
|
||||
pass
|
||||
|
||||
|
||||
class UserInputRequiredException(ToolException):
|
||||
"""Raised when a tool wrapping a sub-agent requires user input to proceed.
|
||||
|
||||
This exception carries the ``user_input_request`` Content items emitted by
|
||||
the sub-agent (e.g., ``oauth_consent_request``, ``function_approval_request``)
|
||||
so the tool invocation layer can propagate them to the parent agent's response
|
||||
instead of swallowing them as a generic tool error.
|
||||
|
||||
Args:
|
||||
contents: The user-input-request Content items from the sub-agent response.
|
||||
message: Human-readable description of why user input is needed.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
contents: list[Any],
|
||||
message: str = "Tool requires user input to proceed.",
|
||||
) -> None:
|
||||
"""Create a UserInputRequiredException.
|
||||
|
||||
Args:
|
||||
contents: The user-input-request Content items from the sub-agent response.
|
||||
message: Human-readable description of why user input is needed.
|
||||
"""
|
||||
super().__init__(message, log_level=None)
|
||||
self.contents = contents
|
||||
|
||||
|
||||
# endregion
|
||||
|
||||
# region Middleware Exceptions
|
||||
|
||||
@@ -1162,11 +1162,35 @@ class ChatTelemetryLayer(Generic[OptionsCoT]):
|
||||
tokenizer: TokenizerProtocol | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Awaitable[ChatResponse[Any]] | ResponseStream[ChatResponseUpdate, ChatResponse[Any]]:
|
||||
"""Trace chat responses with OpenTelemetry spans and metrics."""
|
||||
"""Trace chat responses with OpenTelemetry spans and metrics.
|
||||
|
||||
Args:
|
||||
messages: The message or messages to send to the model.
|
||||
stream: Whether to stream the response. Defaults to False.
|
||||
options: Chat options as a TypedDict.
|
||||
compaction_strategy: Optional compaction strategy to apply before model calls.
|
||||
tokenizer: Optional tokenizer used by token-aware compaction strategies.
|
||||
|
||||
Keyword Args:
|
||||
kwargs: Compatibility keyword arguments from higher client layers. This layer does
|
||||
not consume ``function_invocation_kwargs`` directly; if present, it is ignored
|
||||
because function invocation has already been processed above. If a ``client_kwargs``
|
||||
mapping is present, it is flattened into ordinary keyword arguments for tracing and
|
||||
forwarding so clients that use those values continue to work while clients that
|
||||
ignore extra kwargs remain compatible.
|
||||
"""
|
||||
from ._types import ChatResponse, ChatResponseUpdate, ResponseStream # type: ignore[reportUnusedImport]
|
||||
|
||||
global OBSERVABILITY_SETTINGS
|
||||
super_get_response = super().get_response # type: ignore[misc]
|
||||
compatibility_client_kwargs = kwargs.pop("client_kwargs", None)
|
||||
kwargs.pop("function_invocation_kwargs", None)
|
||||
merged_client_kwargs = (
|
||||
dict(cast(Mapping[str, Any], compatibility_client_kwargs))
|
||||
if isinstance(compatibility_client_kwargs, Mapping)
|
||||
else {}
|
||||
)
|
||||
merged_client_kwargs.update(kwargs)
|
||||
|
||||
if not OBSERVABILITY_SETTINGS.ENABLED:
|
||||
return super_get_response( # type: ignore[no-any-return]
|
||||
@@ -1175,12 +1199,14 @@ class ChatTelemetryLayer(Generic[OptionsCoT]):
|
||||
options=options,
|
||||
compaction_strategy=compaction_strategy,
|
||||
tokenizer=tokenizer,
|
||||
**kwargs,
|
||||
**merged_client_kwargs,
|
||||
)
|
||||
|
||||
opts: dict[str, Any] = options or {} # type: ignore[assignment]
|
||||
provider_name = str(getattr(self, "otel_provider_name", "unknown"))
|
||||
model_id = kwargs.get("model_id") or opts.get("model_id") or getattr(self, "model_id", None) or "unknown"
|
||||
model_id = (
|
||||
merged_client_kwargs.get("model_id") or opts.get("model_id") or getattr(self, "model_id", None) or "unknown"
|
||||
)
|
||||
service_url_func = getattr(self, "service_url", None)
|
||||
service_url = str(service_url_func() if callable(service_url_func) else "unknown")
|
||||
attributes = _get_span_attributes(
|
||||
@@ -1188,7 +1214,7 @@ class ChatTelemetryLayer(Generic[OptionsCoT]):
|
||||
provider_name=provider_name,
|
||||
model=model_id,
|
||||
service_url=service_url,
|
||||
**kwargs,
|
||||
**merged_client_kwargs,
|
||||
)
|
||||
|
||||
if stream:
|
||||
@@ -1200,7 +1226,7 @@ class ChatTelemetryLayer(Generic[OptionsCoT]):
|
||||
options=opts,
|
||||
compaction_strategy=compaction_strategy,
|
||||
tokenizer=tokenizer,
|
||||
**kwargs,
|
||||
**merged_client_kwargs,
|
||||
),
|
||||
)
|
||||
|
||||
@@ -1291,7 +1317,7 @@ class ChatTelemetryLayer(Generic[OptionsCoT]):
|
||||
options=opts,
|
||||
compaction_strategy=compaction_strategy,
|
||||
tokenizer=tokenizer,
|
||||
**kwargs,
|
||||
**merged_client_kwargs,
|
||||
),
|
||||
)
|
||||
except Exception as exception:
|
||||
@@ -1420,6 +1446,8 @@ class AgentTelemetryLayer:
|
||||
session: AgentSession | None = None,
|
||||
compaction_strategy: CompactionStrategy | None = None,
|
||||
tokenizer: TokenizerProtocol | None = None,
|
||||
function_invocation_kwargs: Mapping[str, Any] | None = None,
|
||||
client_kwargs: Mapping[str, Any] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Awaitable[AgentResponse[Any]]: ...
|
||||
|
||||
@@ -1432,6 +1460,8 @@ class AgentTelemetryLayer:
|
||||
session: AgentSession | None = None,
|
||||
compaction_strategy: CompactionStrategy | None = None,
|
||||
tokenizer: TokenizerProtocol | None = None,
|
||||
function_invocation_kwargs: Mapping[str, Any] | None = None,
|
||||
client_kwargs: Mapping[str, Any] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> ResponseStream[AgentResponseUpdate, AgentResponse[Any]]: ...
|
||||
|
||||
@@ -1443,6 +1473,8 @@ class AgentTelemetryLayer:
|
||||
session: AgentSession | None = None,
|
||||
compaction_strategy: CompactionStrategy | None = None,
|
||||
tokenizer: TokenizerProtocol | None = None,
|
||||
function_invocation_kwargs: Mapping[str, Any] | None = None,
|
||||
client_kwargs: Mapping[str, Any] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Awaitable[AgentResponse[Any]] | ResponseStream[AgentResponseUpdate, AgentResponse[Any]]:
|
||||
"""Trace agent runs with OpenTelemetry spans and metrics."""
|
||||
@@ -1463,11 +1495,15 @@ class AgentTelemetryLayer:
|
||||
session=session,
|
||||
compaction_strategy=compaction_strategy,
|
||||
tokenizer=tokenizer,
|
||||
function_invocation_kwargs=function_invocation_kwargs,
|
||||
client_kwargs=client_kwargs,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
default_options = getattr(self, "default_options", {})
|
||||
options = kwargs.get("options")
|
||||
merged_client_kwargs = dict(client_kwargs) if client_kwargs is not None else {}
|
||||
merged_client_kwargs.update(kwargs)
|
||||
merged_options: dict[str, Any] = merge_chat_options(default_options, options or {})
|
||||
attributes = _get_span_attributes(
|
||||
operation_name=OtelAttr.AGENT_INVOKE_OPERATION,
|
||||
@@ -1477,7 +1513,7 @@ class AgentTelemetryLayer:
|
||||
agent_description=getattr(self, "description", None),
|
||||
thread_id=session.service_session_id if session else None,
|
||||
all_options=merged_options,
|
||||
**kwargs,
|
||||
**merged_client_kwargs,
|
||||
)
|
||||
|
||||
if stream:
|
||||
@@ -1487,6 +1523,8 @@ class AgentTelemetryLayer:
|
||||
session=session,
|
||||
compaction_strategy=compaction_strategy,
|
||||
tokenizer=tokenizer,
|
||||
function_invocation_kwargs=function_invocation_kwargs,
|
||||
client_kwargs=client_kwargs,
|
||||
**kwargs,
|
||||
)
|
||||
if isinstance(run_result, ResponseStream):
|
||||
@@ -1578,6 +1616,8 @@ class AgentTelemetryLayer:
|
||||
session=session,
|
||||
compaction_strategy=compaction_strategy,
|
||||
tokenizer=tokenizer,
|
||||
function_invocation_kwargs=function_invocation_kwargs,
|
||||
client_kwargs=client_kwargs,
|
||||
**kwargs,
|
||||
)
|
||||
except Exception as exception:
|
||||
|
||||
@@ -15,7 +15,7 @@ from collections.abc import (
|
||||
)
|
||||
from datetime import datetime, timezone
|
||||
from itertools import chain
|
||||
from typing import Any, Generic, Literal, cast
|
||||
from typing import Any, Generic, Literal, cast, overload
|
||||
|
||||
from openai import AsyncOpenAI, BadRequestError
|
||||
from openai.lib._parsing._completions import type_to_response_format_param
|
||||
@@ -30,7 +30,8 @@ from openai.types.chat.completion_create_params import WebSearchOptions
|
||||
from pydantic import BaseModel
|
||||
|
||||
from .._clients import BaseChatClient
|
||||
from .._middleware import ChatAndFunctionMiddlewareTypes, ChatMiddlewareLayer
|
||||
from .._docstrings import apply_layered_docstring
|
||||
from .._middleware import ChatAndFunctionMiddlewareTypes, ChatMiddlewareLayer, FunctionMiddlewareTypes
|
||||
from .._settings import load_settings
|
||||
from .._tools import (
|
||||
FunctionInvocationConfiguration,
|
||||
@@ -72,6 +73,7 @@ else:
|
||||
|
||||
logger = logging.getLogger("agent_framework.openai")
|
||||
|
||||
ResponseModelBoundT = TypeVar("ResponseModelBoundT", bound=BaseModel)
|
||||
ResponseModelT = TypeVar("ResponseModelT", bound=BaseModel | None, default=None)
|
||||
|
||||
|
||||
@@ -213,6 +215,57 @@ class RawOpenAIChatClient( # type: ignore[misc]
|
||||
|
||||
# endregion
|
||||
|
||||
@overload
|
||||
def get_response(
|
||||
self,
|
||||
messages: Sequence[Message],
|
||||
*,
|
||||
stream: Literal[False] = ...,
|
||||
options: ChatOptions[ResponseModelBoundT],
|
||||
**kwargs: Any,
|
||||
) -> Awaitable[ChatResponse[ResponseModelBoundT]]: ...
|
||||
|
||||
@overload
|
||||
def get_response(
|
||||
self,
|
||||
messages: Sequence[Message],
|
||||
*,
|
||||
stream: Literal[False] = ...,
|
||||
options: OpenAIChatOptionsT | ChatOptions[None] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Awaitable[ChatResponse[Any]]: ...
|
||||
|
||||
@overload
|
||||
def get_response(
|
||||
self,
|
||||
messages: Sequence[Message],
|
||||
*,
|
||||
stream: Literal[True],
|
||||
options: OpenAIChatOptionsT | ChatOptions[Any] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> ResponseStream[ChatResponseUpdate, ChatResponse[Any]]: ...
|
||||
|
||||
@override
|
||||
def get_response(
|
||||
self,
|
||||
messages: Sequence[Message],
|
||||
*,
|
||||
stream: bool = False,
|
||||
options: OpenAIChatOptionsT | ChatOptions[Any] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Awaitable[ChatResponse[Any]] | ResponseStream[ChatResponseUpdate, ChatResponse[Any]]:
|
||||
"""Get a response from the raw OpenAI chat client."""
|
||||
super_get_response = cast(
|
||||
"Callable[..., Awaitable[ChatResponse[Any]] | ResponseStream[ChatResponseUpdate, ChatResponse[Any]]]",
|
||||
super().get_response, # type: ignore[misc]
|
||||
)
|
||||
return super_get_response( # type: ignore[no-any-return]
|
||||
messages=messages,
|
||||
stream=stream,
|
||||
options=options,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
@override
|
||||
def _inner_get_response(
|
||||
self,
|
||||
@@ -727,6 +780,77 @@ class OpenAIChatClient( # type: ignore[misc]
|
||||
):
|
||||
"""OpenAI Chat completion class with middleware, telemetry, and function invocation support."""
|
||||
|
||||
@overload
|
||||
def get_response(
|
||||
self,
|
||||
messages: Sequence[Message],
|
||||
*,
|
||||
stream: Literal[False] = ...,
|
||||
options: ChatOptions[ResponseModelBoundT],
|
||||
function_middleware: Sequence[FunctionMiddlewareTypes] | None = None,
|
||||
function_invocation_kwargs: Mapping[str, Any] | None = None,
|
||||
client_kwargs: Mapping[str, Any] | None = None,
|
||||
middleware: Sequence[ChatAndFunctionMiddlewareTypes] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Awaitable[ChatResponse[ResponseModelBoundT]]: ...
|
||||
|
||||
@overload
|
||||
def get_response(
|
||||
self,
|
||||
messages: Sequence[Message],
|
||||
*,
|
||||
stream: Literal[False] = ...,
|
||||
options: OpenAIChatOptionsT | ChatOptions[None] | None = None,
|
||||
function_middleware: Sequence[FunctionMiddlewareTypes] | None = None,
|
||||
function_invocation_kwargs: Mapping[str, Any] | None = None,
|
||||
client_kwargs: Mapping[str, Any] | None = None,
|
||||
middleware: Sequence[ChatAndFunctionMiddlewareTypes] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Awaitable[ChatResponse[Any]]: ...
|
||||
|
||||
@overload
|
||||
def get_response(
|
||||
self,
|
||||
messages: Sequence[Message],
|
||||
*,
|
||||
stream: Literal[True],
|
||||
options: OpenAIChatOptionsT | ChatOptions[Any] | None = None,
|
||||
function_middleware: Sequence[FunctionMiddlewareTypes] | None = None,
|
||||
function_invocation_kwargs: Mapping[str, Any] | None = None,
|
||||
client_kwargs: Mapping[str, Any] | None = None,
|
||||
middleware: Sequence[ChatAndFunctionMiddlewareTypes] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> ResponseStream[ChatResponseUpdate, ChatResponse[Any]]: ...
|
||||
|
||||
@override
|
||||
def get_response(
|
||||
self,
|
||||
messages: Sequence[Message],
|
||||
*,
|
||||
stream: bool = False,
|
||||
options: OpenAIChatOptionsT | ChatOptions[Any] | None = None,
|
||||
function_middleware: Sequence[FunctionMiddlewareTypes] | None = None,
|
||||
function_invocation_kwargs: Mapping[str, Any] | None = None,
|
||||
client_kwargs: Mapping[str, Any] | None = None,
|
||||
middleware: Sequence[ChatAndFunctionMiddlewareTypes] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Awaitable[ChatResponse[Any]] | ResponseStream[ChatResponseUpdate, ChatResponse[Any]]:
|
||||
"""Get a response from the OpenAI chat client with all standard layers enabled."""
|
||||
super_get_response = cast(
|
||||
"Callable[..., Awaitable[ChatResponse[Any]] | ResponseStream[ChatResponseUpdate, ChatResponse[Any]]]",
|
||||
super().get_response, # type: ignore[misc]
|
||||
)
|
||||
return super_get_response( # type: ignore[no-any-return]
|
||||
messages=messages,
|
||||
stream=stream,
|
||||
options=options,
|
||||
function_middleware=function_middleware,
|
||||
function_invocation_kwargs=function_invocation_kwargs,
|
||||
client_kwargs=client_kwargs,
|
||||
middleware=middleware,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
@@ -830,3 +954,25 @@ class OpenAIChatClient( # type: ignore[misc]
|
||||
middleware=middleware,
|
||||
function_invocation_configuration=function_invocation_configuration,
|
||||
)
|
||||
|
||||
|
||||
def _apply_openai_chat_client_docstrings() -> None:
|
||||
"""Align OpenAI chat-client docstrings with the raw implementation."""
|
||||
apply_layered_docstring(RawOpenAIChatClient.get_response, BaseChatClient.get_response)
|
||||
apply_layered_docstring(
|
||||
OpenAIChatClient.get_response,
|
||||
RawOpenAIChatClient.get_response,
|
||||
extra_keyword_args={
|
||||
"function_middleware": """
|
||||
Optional per-call function middleware.
|
||||
When omitted, middleware configured on the client or forwarded from higher layers is used.
|
||||
""",
|
||||
"middleware": """
|
||||
Optional per-call chat and function middleware.
|
||||
This is merged with any middleware configured on the client for the current request.
|
||||
""",
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
_apply_openai_chat_client_docstrings()
|
||||
|
||||
@@ -665,7 +665,7 @@ class RawOpenAIResponsesClient( # type: ignore[misc]
|
||||
if output_format:
|
||||
tool["output_format"] = output_format
|
||||
if model:
|
||||
tool["model"] = model
|
||||
tool["model"] = model # type: ignore
|
||||
if quality:
|
||||
tool["quality"] = quality
|
||||
if partial_images is not None:
|
||||
|
||||
@@ -24,19 +24,19 @@ classifiers = [
|
||||
]
|
||||
dependencies = [
|
||||
# utilities
|
||||
"typing-extensions",
|
||||
"typing-extensions>=4.15.0,<5",
|
||||
"pydantic>=2,<3",
|
||||
"python-dotenv>=1,<2",
|
||||
# telemetry
|
||||
"opentelemetry-api>=1.39.0",
|
||||
"opentelemetry-sdk>=1.39.0",
|
||||
"opentelemetry-semantic-conventions-ai>=0.4.13",
|
||||
"opentelemetry-api>=1.39.0,<2",
|
||||
"opentelemetry-sdk>=1.39.0,<2",
|
||||
"opentelemetry-semantic-conventions-ai>=0.4.13,<0.4.14",
|
||||
# connectors and functions
|
||||
"openai>=1.99.0",
|
||||
"openai>=1.99.0,<3",
|
||||
"azure-identity>=1,<2",
|
||||
"azure-ai-projects>=2.0.0,<3.0",
|
||||
"mcp[ws]>=1.24.0,<2",
|
||||
"packaging>=24.1",
|
||||
"packaging>=24.1,<25",
|
||||
]
|
||||
|
||||
[project.optional-dependencies]
|
||||
@@ -76,15 +76,7 @@ environments = [
|
||||
fallback-version = "0.0.0"
|
||||
|
||||
[tool.pytest.ini_options]
|
||||
testpaths = [
|
||||
'tests',
|
||||
'packages/core/tests',
|
||||
'packages/a2a/tests',
|
||||
'packages/azure-ai/tests',
|
||||
'packages/copilotstudio/tests',
|
||||
'packages/mem0/tests',
|
||||
'packages/runtime/tests'
|
||||
]
|
||||
testpaths = ['tests']
|
||||
addopts = "-ra -q -r fEX"
|
||||
asyncio_mode = "auto"
|
||||
asyncio_default_fixture_loop_scope = "function"
|
||||
@@ -131,7 +123,7 @@ include = "../../shared_tasks.toml"
|
||||
|
||||
[tool.poe.tasks]
|
||||
mypy = "mypy --config-file $POE_ROOT/pyproject.toml agent_framework"
|
||||
test = "pytest -m \"not integration\" --cov=agent_framework --cov-report=term-missing:skip-covered -n auto --dist worksteal tests"
|
||||
test = 'pytest -m "not integration" --cov=agent_framework --cov-report=term-missing:skip-covered -n auto --dist worksteal tests'
|
||||
|
||||
[tool.flit.module]
|
||||
name = "agent_framework"
|
||||
|
||||
@@ -0,0 +1,159 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
from openai.types import CreateEmbeddingResponse
|
||||
from openai.types import Embedding as OpenAIEmbedding
|
||||
from openai.types.create_embedding_response import Usage
|
||||
|
||||
from agent_framework.azure import AzureOpenAIEmbeddingClient
|
||||
from agent_framework.openai import OpenAIEmbeddingOptions
|
||||
|
||||
|
||||
def _make_openai_response(
|
||||
embeddings: list[list[float]],
|
||||
model: str = "text-embedding-3-small",
|
||||
prompt_tokens: int = 5,
|
||||
total_tokens: int = 5,
|
||||
) -> CreateEmbeddingResponse:
|
||||
"""Helper to create a mock OpenAI embeddings response."""
|
||||
data = [OpenAIEmbedding(embedding=emb, index=i, object="embedding") for i, emb in enumerate(embeddings)]
|
||||
return CreateEmbeddingResponse(
|
||||
data=data,
|
||||
model=model,
|
||||
object="list",
|
||||
usage=Usage(prompt_tokens=prompt_tokens, total_tokens=total_tokens),
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def azure_embedding_unit_test_env(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
"""Clear ambient Azure OpenAI embedding env vars for deterministic unit tests."""
|
||||
for key in (
|
||||
"AZURE_OPENAI_ENDPOINT",
|
||||
"AZURE_OPENAI_API_KEY",
|
||||
"AZURE_OPENAI_EMBEDDING_DEPLOYMENT_NAME",
|
||||
"AZURE_OPENAI_BASE_URL",
|
||||
"AZURE_OPENAI_TOKEN_ENDPOINT",
|
||||
):
|
||||
monkeypatch.delenv(key, raising=False)
|
||||
|
||||
|
||||
def test_azure_construction_with_deployment_name(azure_embedding_unit_test_env: None) -> None:
|
||||
client = AzureOpenAIEmbeddingClient(
|
||||
deployment_name="text-embedding-3-small",
|
||||
api_key="test-key",
|
||||
endpoint="https://test.openai.azure.com/",
|
||||
)
|
||||
assert client.model_id == "text-embedding-3-small"
|
||||
|
||||
|
||||
def test_azure_construction_with_existing_client(azure_embedding_unit_test_env: None) -> None:
|
||||
mock_client = MagicMock()
|
||||
client = AzureOpenAIEmbeddingClient(
|
||||
deployment_name="my-deployment",
|
||||
async_client=mock_client,
|
||||
)
|
||||
assert client.model_id == "my-deployment"
|
||||
assert client.client is mock_client
|
||||
|
||||
|
||||
def test_azure_construction_missing_deployment_name_raises(azure_embedding_unit_test_env: None) -> None:
|
||||
with pytest.raises(ValueError, match="deployment name is required"):
|
||||
AzureOpenAIEmbeddingClient(
|
||||
api_key="test-key",
|
||||
endpoint="https://test.openai.azure.com/",
|
||||
)
|
||||
|
||||
|
||||
def test_azure_construction_missing_credentials_raises(azure_embedding_unit_test_env: None) -> None:
|
||||
with pytest.raises(ValueError, match="api_key, credential, or a client"):
|
||||
AzureOpenAIEmbeddingClient(
|
||||
deployment_name="test",
|
||||
endpoint="https://test.openai.azure.com/",
|
||||
)
|
||||
|
||||
|
||||
async def test_azure_get_embeddings(azure_embedding_unit_test_env: None) -> None:
|
||||
mock_response = _make_openai_response(
|
||||
embeddings=[[0.1, 0.2]],
|
||||
)
|
||||
mock_async_client = MagicMock()
|
||||
mock_async_client.embeddings = MagicMock()
|
||||
mock_async_client.embeddings.create = AsyncMock(return_value=mock_response)
|
||||
|
||||
client = AzureOpenAIEmbeddingClient(
|
||||
deployment_name="text-embedding-3-small",
|
||||
async_client=mock_async_client,
|
||||
)
|
||||
|
||||
result = await client.get_embeddings(["hello"])
|
||||
|
||||
assert len(result) == 1
|
||||
assert result[0].vector == [0.1, 0.2]
|
||||
|
||||
|
||||
def test_azure_otel_provider_name(azure_embedding_unit_test_env: None) -> None:
|
||||
mock_client = MagicMock()
|
||||
client = AzureOpenAIEmbeddingClient(
|
||||
deployment_name="test",
|
||||
async_client=mock_client,
|
||||
)
|
||||
assert client.OTEL_PROVIDER_NAME == "azure.ai.openai"
|
||||
|
||||
|
||||
skip_if_azure_openai_integration_tests_disabled = pytest.mark.skipif(
|
||||
not os.getenv("AZURE_OPENAI_ENDPOINT")
|
||||
or (not os.getenv("AZURE_OPENAI_API_KEY") and not os.getenv("AZURE_OPENAI_EMBEDDING_DEPLOYMENT_NAME")),
|
||||
reason="No Azure OpenAI credentials provided; skipping integration tests.",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.flaky
|
||||
@pytest.mark.integration
|
||||
@skip_if_azure_openai_integration_tests_disabled
|
||||
async def test_integration_azure_openai_get_embeddings() -> None:
|
||||
"""End-to-end test of Azure OpenAI embedding generation."""
|
||||
client = AzureOpenAIEmbeddingClient()
|
||||
|
||||
result = await client.get_embeddings(["hello world"])
|
||||
|
||||
assert len(result) == 1
|
||||
assert isinstance(result[0].vector, list)
|
||||
assert len(result[0].vector) > 0
|
||||
assert all(isinstance(v, float) for v in result[0].vector)
|
||||
assert result[0].model_id is not None
|
||||
assert result.usage is not None
|
||||
assert result.usage["input_token_count"] > 0
|
||||
|
||||
|
||||
@pytest.mark.flaky
|
||||
@pytest.mark.integration
|
||||
@skip_if_azure_openai_integration_tests_disabled
|
||||
async def test_integration_azure_openai_get_embeddings_multiple() -> None:
|
||||
"""Test Azure OpenAI embedding generation for multiple inputs."""
|
||||
client = AzureOpenAIEmbeddingClient()
|
||||
|
||||
result = await client.get_embeddings(["hello", "world", "test"])
|
||||
|
||||
assert len(result) == 3
|
||||
dims = [len(e.vector) for e in result]
|
||||
assert all(d == dims[0] for d in dims)
|
||||
|
||||
|
||||
@pytest.mark.flaky
|
||||
@pytest.mark.integration
|
||||
@skip_if_azure_openai_integration_tests_disabled
|
||||
async def test_integration_azure_openai_get_embeddings_with_dimensions() -> None:
|
||||
"""Test Azure OpenAI embedding generation with custom dimensions."""
|
||||
client = AzureOpenAIEmbeddingClient()
|
||||
|
||||
options: OpenAIEmbeddingOptions = {"dimensions": 256}
|
||||
result = await client.get_embeddings(["hello world"], options=options)
|
||||
|
||||
assert len(result) == 1
|
||||
assert len(result[0].vector) == 256
|
||||
@@ -1,6 +1,7 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
import contextlib
|
||||
import inspect
|
||||
from collections.abc import AsyncIterable, MutableSequence
|
||||
from typing import Any
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
@@ -30,7 +31,8 @@ from agent_framework import (
|
||||
tool,
|
||||
)
|
||||
from agent_framework._agents import _get_tool_name, _merge_options, _sanitize_agent_name
|
||||
from agent_framework._mcp import MCPTool
|
||||
from agent_framework._mcp import MCPTool, _build_prefixed_mcp_name, _normalize_mcp_name
|
||||
from agent_framework._middleware import FunctionInvocationContext
|
||||
|
||||
|
||||
class _FixedTokenizer:
|
||||
@@ -41,6 +43,30 @@ class _FixedTokenizer:
|
||||
return self.token_count
|
||||
|
||||
|
||||
class _ConnectedMCPTool(MCPTool):
|
||||
def __init__(self, name: str, function_names: list[str], *, tool_name_prefix: str | None = None) -> None:
|
||||
super().__init__(name=name, tool_name_prefix=tool_name_prefix)
|
||||
self.is_connected = True
|
||||
self._functions = []
|
||||
for function_name in function_names:
|
||||
normalized_name = _normalize_mcp_name(function_name)
|
||||
exposed_name = _build_prefixed_mcp_name(normalized_name, self.tool_name_prefix)
|
||||
self._functions.append(
|
||||
FunctionTool(
|
||||
func=lambda value=function_name: value,
|
||||
name=exposed_name,
|
||||
description=f"{function_name} from {name}",
|
||||
additional_properties={
|
||||
"_mcp_remote_name": function_name,
|
||||
"_mcp_normalized_name": normalized_name,
|
||||
},
|
||||
)
|
||||
)
|
||||
|
||||
def get_mcp_client(self) -> contextlib.AbstractAsyncContextManager[Any]:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
def test_agent_session_type(agent_session: AgentSession) -> None:
|
||||
assert isinstance(agent_session, AgentSession)
|
||||
|
||||
@@ -77,6 +103,30 @@ def test_chat_client_agent_type(client: SupportsChatGetResponse) -> None:
|
||||
assert isinstance(chat_client_agent, SupportsAgentRun)
|
||||
|
||||
|
||||
def test_agent_init_docstring_surfaces_raw_agent_constructor_docs() -> None:
|
||||
docstring = inspect.getdoc(Agent.__init__)
|
||||
|
||||
assert docstring is not None
|
||||
assert "client: The chat client to use for the agent." in docstring
|
||||
assert "middleware: List of middleware to intercept agent and function invocations." in docstring
|
||||
|
||||
|
||||
def test_agent_run_docstring_surfaces_raw_agent_runtime_docs() -> None:
|
||||
docstring = inspect.getdoc(Agent.run)
|
||||
|
||||
assert docstring is not None
|
||||
assert "Run the agent with the given messages and options." in docstring
|
||||
assert "function_invocation_kwargs: Keyword arguments forwarded to tool invocation." in docstring
|
||||
assert "middleware: Optional per-run agent, chat, and function middleware." in docstring
|
||||
|
||||
|
||||
def test_agent_run_is_defined_on_agent_class() -> None:
|
||||
signature = inspect.signature(Agent.run)
|
||||
|
||||
assert Agent.run.__qualname__ == "Agent.run"
|
||||
assert "middleware" in signature.parameters
|
||||
|
||||
|
||||
async def test_chat_client_agent_init(client: SupportsChatGetResponse) -> None:
|
||||
agent_id = str(uuid4())
|
||||
agent = Agent(client=client, id=agent_id, description="Test")
|
||||
@@ -97,6 +147,13 @@ async def test_chat_client_agent_init_with_name(
|
||||
assert agent.description == "Test"
|
||||
|
||||
|
||||
def test_agent_init_warns_for_direct_additional_properties(client: SupportsChatGetResponse) -> None:
|
||||
with pytest.warns(DeprecationWarning, match="additional_properties"):
|
||||
agent = Agent(client=client, legacy_key="legacy-value")
|
||||
|
||||
assert agent.additional_properties["legacy_key"] == "legacy-value"
|
||||
|
||||
|
||||
async def test_chat_client_agent_run(client: SupportsChatGetResponse) -> None:
|
||||
agent = Agent(client=client)
|
||||
|
||||
@@ -229,33 +286,38 @@ async def test_prepare_session_does_not_mutate_agent_chat_options(
|
||||
assert len(agent.default_options["tools"]) == 1
|
||||
|
||||
|
||||
async def test_prepare_run_context_keeps_compaction_overrides_out_of_kwargs(
|
||||
async def test_prepare_run_context_handles_function_kwargs(
|
||||
chat_client_base: SupportsChatGetResponse,
|
||||
) -> None:
|
||||
strategy = SlidingWindowStrategy(keep_last_groups=2)
|
||||
tokenizer = _FixedTokenizer(13)
|
||||
agent = Agent(client=chat_client_base)
|
||||
session = agent.create_session()
|
||||
|
||||
ctx = await agent._prepare_run_context( # type: ignore[reportPrivateUsage]
|
||||
messages=[Message(role="user", text="Hello")],
|
||||
session=None,
|
||||
messages="Hello",
|
||||
session=session,
|
||||
tools=None,
|
||||
options=None,
|
||||
compaction_strategy=strategy,
|
||||
tokenizer=tokenizer,
|
||||
kwargs={"custom_flag": True},
|
||||
options={
|
||||
"temperature": 0.4,
|
||||
"additional_function_arguments": {"from_options": "options-value"},
|
||||
},
|
||||
compaction_strategy=None,
|
||||
tokenizer=None,
|
||||
legacy_kwargs={"legacy_key": "legacy-value"},
|
||||
function_invocation_kwargs={"runtime_key": "runtime-value"},
|
||||
client_kwargs={"client_key": "client-value"},
|
||||
)
|
||||
|
||||
assert ctx["compaction_strategy"] is strategy
|
||||
assert ctx["tokenizer"] is tokenizer
|
||||
assert ctx["filtered_kwargs"].get("custom_flag") is True
|
||||
assert "compaction_strategy" not in ctx["filtered_kwargs"]
|
||||
assert "tokenizer" not in ctx["filtered_kwargs"]
|
||||
assert ctx["chat_options"]["temperature"] == 0.4
|
||||
assert "additional_function_arguments" not in ctx["chat_options"]
|
||||
assert ctx["function_invocation_kwargs"]["from_options"] == "options-value"
|
||||
assert ctx["function_invocation_kwargs"]["legacy_key"] == "legacy-value"
|
||||
assert ctx["function_invocation_kwargs"]["runtime_key"] == "runtime-value"
|
||||
assert "session" not in ctx["function_invocation_kwargs"]
|
||||
assert ctx["client_kwargs"]["client_key"] == "client-value"
|
||||
assert ctx["client_kwargs"]["session"] is session
|
||||
|
||||
|
||||
async def test_chat_client_agent_run_with_session(
|
||||
chat_client_base: SupportsChatGetResponse,
|
||||
) -> None:
|
||||
async def test_chat_client_agent_run_with_session(chat_client_base: SupportsChatGetResponse) -> None:
|
||||
mock_response = ChatResponse(
|
||||
messages=[Message(role="assistant", contents=[Content.from_text("test response")])],
|
||||
conversation_id="123",
|
||||
@@ -696,8 +758,9 @@ async def test_chat_agent_as_tool_basic(client: SupportsChatGetResponse) -> None
|
||||
|
||||
assert tool.name == "TestAgent"
|
||||
assert tool.description == "Test agent for as_tool"
|
||||
assert tool.approval_mode == "never_require"
|
||||
assert hasattr(tool, "func")
|
||||
assert hasattr(tool, "input_model")
|
||||
assert tool.input_model is None
|
||||
|
||||
|
||||
async def test_chat_agent_as_tool_custom_parameters(
|
||||
@@ -711,13 +774,15 @@ async def test_chat_agent_as_tool_custom_parameters(
|
||||
description="Custom description",
|
||||
arg_name="query",
|
||||
arg_description="Custom input description",
|
||||
approval_mode="always_require",
|
||||
)
|
||||
|
||||
assert tool.name == "CustomTool"
|
||||
assert tool.description == "Custom description"
|
||||
assert tool.approval_mode == "always_require"
|
||||
|
||||
# Check that the input model has the custom field name
|
||||
schema = tool.input_model.model_json_schema()
|
||||
schema = tool.parameters()
|
||||
assert "query" in schema["properties"]
|
||||
assert schema["properties"]["query"]["description"] == "Custom input description"
|
||||
|
||||
@@ -736,7 +801,7 @@ async def test_chat_agent_as_tool_defaults(client: SupportsChatGetResponse) -> N
|
||||
assert tool.description == "" # Should default to empty string
|
||||
|
||||
# Check default input field
|
||||
schema = tool.input_model.model_json_schema()
|
||||
schema = tool.parameters()
|
||||
assert "task" in schema["properties"]
|
||||
assert "Task for TestAgent" in schema["properties"]["task"]["description"]
|
||||
|
||||
@@ -759,12 +824,12 @@ async def test_chat_agent_as_tool_function_execution(
|
||||
tool = agent.as_tool()
|
||||
|
||||
# Test function execution
|
||||
result = await tool.invoke(arguments=tool.input_model(task="Hello"))
|
||||
result = await tool.invoke(arguments={"task": "Hello"})
|
||||
|
||||
# Should return the agent's response text as a list of Content items
|
||||
assert isinstance(result, list)
|
||||
assert len(result) == 1
|
||||
assert result[0].text == "test response" # From mock chat client
|
||||
assert result[0].text == "test streaming response another update" # From mock streaming client
|
||||
|
||||
|
||||
async def test_chat_agent_as_tool_with_stream_callback(
|
||||
@@ -782,7 +847,7 @@ async def test_chat_agent_as_tool_with_stream_callback(
|
||||
tool = agent.as_tool(stream_callback=stream_callback)
|
||||
|
||||
# Execute the tool
|
||||
result = await tool.invoke(arguments=tool.input_model(task="Hello"))
|
||||
result = await tool.invoke(arguments={"task": "Hello"})
|
||||
|
||||
# Should have collected streaming updates
|
||||
assert len(collected_updates) > 0
|
||||
@@ -802,9 +867,9 @@ async def test_chat_agent_as_tool_with_custom_arg_name(
|
||||
tool = agent.as_tool(arg_name="prompt", arg_description="Custom prompt input")
|
||||
|
||||
# Test that the custom argument name works
|
||||
result = await tool.invoke(arguments=tool.input_model(prompt="Test prompt"))
|
||||
result = await tool.invoke(arguments={"prompt": "Test prompt"})
|
||||
assert isinstance(result, list)
|
||||
assert result[0].text == "test response"
|
||||
assert result[0].text == "test streaming response another update"
|
||||
|
||||
|
||||
async def test_chat_agent_as_tool_with_async_stream_callback(
|
||||
@@ -822,7 +887,7 @@ async def test_chat_agent_as_tool_with_async_stream_callback(
|
||||
tool = agent.as_tool(stream_callback=async_stream_callback)
|
||||
|
||||
# Execute the tool
|
||||
result = await tool.invoke(arguments=tool.input_model(task="Hello"))
|
||||
result = await tool.invoke(arguments={"task": "Hello"})
|
||||
|
||||
# Should have collected streaming updates
|
||||
assert len(collected_updates) > 0
|
||||
@@ -853,17 +918,14 @@ async def test_chat_agent_as_tool_name_sanitization(
|
||||
assert tool.name == expected_tool_name, f"Expected {expected_tool_name}, got {tool.name} for input {agent_name}"
|
||||
|
||||
|
||||
async def test_chat_agent_as_tool_propagate_session_true(
|
||||
client: SupportsChatGetResponse,
|
||||
) -> None:
|
||||
"""Test that propagate_session=True forwards the parent's session to the sub-agent."""
|
||||
async def test_chat_agent_as_tool_propagate_session_true(client: SupportsChatGetResponse) -> None:
|
||||
"""Test that propagate_session=True forwards the session to the sub-agent."""
|
||||
agent = Agent(client=client, name="SubAgent", description="Sub agent")
|
||||
tool = agent.as_tool(propagate_session=True)
|
||||
|
||||
parent_session = AgentSession(session_id="parent-session-123")
|
||||
parent_session.state["shared_key"] = "shared_value"
|
||||
|
||||
# Spy on the agent's run method to capture the session argument
|
||||
original_run = agent.run
|
||||
captured_session = None
|
||||
|
||||
@@ -874,16 +936,20 @@ async def test_chat_agent_as_tool_propagate_session_true(
|
||||
|
||||
agent.run = capturing_run # type: ignore[assignment, method-assign]
|
||||
|
||||
await tool.invoke(arguments=tool.input_model(task="Hello"), session=parent_session)
|
||||
await tool.invoke(
|
||||
context=FunctionInvocationContext(
|
||||
function=tool,
|
||||
arguments={"task": "Hello"},
|
||||
session=parent_session,
|
||||
)
|
||||
)
|
||||
|
||||
assert captured_session is parent_session
|
||||
assert captured_session.session_id == "parent-session-123"
|
||||
assert captured_session.state["shared_key"] == "shared_value"
|
||||
|
||||
|
||||
async def test_chat_agent_as_tool_propagate_session_false_by_default(
|
||||
client: SupportsChatGetResponse,
|
||||
) -> None:
|
||||
async def test_chat_agent_as_tool_propagate_session_false_by_default(client: SupportsChatGetResponse) -> None:
|
||||
"""Test that propagate_session defaults to False and does not forward the session."""
|
||||
agent = Agent(client=client, name="SubAgent", description="Sub agent")
|
||||
tool = agent.as_tool() # default: propagate_session=False
|
||||
@@ -900,22 +966,25 @@ async def test_chat_agent_as_tool_propagate_session_false_by_default(
|
||||
|
||||
agent.run = capturing_run # type: ignore[assignment, method-assign]
|
||||
|
||||
await tool.invoke(arguments=tool.input_model(task="Hello"), session=parent_session)
|
||||
await tool.invoke(
|
||||
context=FunctionInvocationContext(
|
||||
function=tool,
|
||||
arguments={"task": "Hello"},
|
||||
session=parent_session,
|
||||
)
|
||||
)
|
||||
|
||||
assert captured_session is None
|
||||
|
||||
|
||||
async def test_chat_agent_as_tool_propagate_session_shares_state(
|
||||
client: SupportsChatGetResponse,
|
||||
) -> None:
|
||||
"""Test that shared session allows the sub-agent to read and write parent's state."""
|
||||
async def test_chat_agent_as_tool_propagate_session_shares_state(client: SupportsChatGetResponse) -> None:
|
||||
"""Test that a propagated session allows the sub-agent to read and write parent state."""
|
||||
agent = Agent(client=client, name="SubAgent", description="Sub agent")
|
||||
tool = agent.as_tool(propagate_session=True)
|
||||
|
||||
parent_session = AgentSession(session_id="shared-session")
|
||||
parent_session.state["counter"] = 0
|
||||
|
||||
# The sub-agent receives the same session object, so mutations are shared
|
||||
original_run = agent.run
|
||||
captured_session = None
|
||||
|
||||
@@ -928,9 +997,14 @@ async def test_chat_agent_as_tool_propagate_session_shares_state(
|
||||
|
||||
agent.run = capturing_run # type: ignore[assignment, method-assign]
|
||||
|
||||
await tool.invoke(arguments=tool.input_model(task="Hello"), session=parent_session)
|
||||
await tool.invoke(
|
||||
context=FunctionInvocationContext(
|
||||
function=tool,
|
||||
arguments={"task": "Hello"},
|
||||
session=parent_session,
|
||||
)
|
||||
)
|
||||
|
||||
# The parent's state should reflect the sub-agent's mutation
|
||||
assert parent_session.state["counter"] == 1
|
||||
|
||||
|
||||
@@ -953,6 +1027,7 @@ async def test_chat_agent_run_with_mcp_tools(client: SupportsChatGetResponse) ->
|
||||
|
||||
# Create a mock MCP tool
|
||||
mock_mcp_tool = MagicMock(spec=MCPTool)
|
||||
mock_mcp_tool.name = "mock-mcp"
|
||||
mock_mcp_tool.is_connected = False
|
||||
mock_mcp_tool.functions = [MagicMock()]
|
||||
|
||||
@@ -970,6 +1045,7 @@ async def test_chat_agent_with_local_mcp_tools(client: SupportsChatGetResponse)
|
||||
"""Test agent initialization with local MCP tools."""
|
||||
# Create a mock MCP tool
|
||||
mock_mcp_tool = MagicMock(spec=MCPTool)
|
||||
mock_mcp_tool.name = "mock-mcp"
|
||||
mock_mcp_tool.is_connected = False
|
||||
mock_mcp_tool.__aenter__ = AsyncMock(return_value=mock_mcp_tool)
|
||||
mock_mcp_tool.__aexit__ = AsyncMock(return_value=None)
|
||||
@@ -1009,6 +1085,7 @@ async def test_mcp_tools_not_duplicated_when_passed_as_runtime_tools(
|
||||
|
||||
# Create a mock MCP tool that is already connected (simulates turn 2)
|
||||
mock_mcp_tool = MagicMock(spec=MCPTool)
|
||||
mock_mcp_tool.name = "mock-mcp"
|
||||
mock_mcp_tool.is_connected = True
|
||||
mock_mcp_tool.functions = [mcp_func_a, mcp_func_b]
|
||||
mock_mcp_tool.__aenter__ = AsyncMock(return_value=mock_mcp_tool)
|
||||
@@ -1032,8 +1109,79 @@ async def test_mcp_tools_not_duplicated_when_passed_as_runtime_tools(
|
||||
assert len(tool_names) == 3
|
||||
|
||||
|
||||
async def test_agent_run_raises_on_local_and_agent_mcp_name_conflict(chat_client_base: Any) -> None:
|
||||
local_tool = FunctionTool(
|
||||
func=lambda: "local",
|
||||
name="delete_all_data",
|
||||
description="Local protected tool",
|
||||
approval_mode="always_require",
|
||||
)
|
||||
agent = Agent(
|
||||
client=chat_client_base,
|
||||
name="TestAgent",
|
||||
tools=[_ConnectedMCPTool(name="dangerous-mcp", function_names=["delete_all_data"])],
|
||||
)
|
||||
|
||||
with raises(ValueError, match="tool_name_prefix"):
|
||||
await agent.run("hello", tools=[local_tool])
|
||||
|
||||
|
||||
async def test_agent_run_raises_on_runtime_local_and_runtime_mcp_name_conflict(chat_client_base: Any) -> None:
|
||||
local_tool = FunctionTool(
|
||||
func=lambda: "local",
|
||||
name="delete_all_data",
|
||||
description="Local protected tool",
|
||||
approval_mode="always_require",
|
||||
)
|
||||
runtime_mcp = _ConnectedMCPTool(name="dangerous-mcp", function_names=["delete_all_data"])
|
||||
agent = Agent(client=chat_client_base, name="TestAgent")
|
||||
|
||||
with raises(ValueError, match="tool_name_prefix"):
|
||||
await agent.run("hello", tools=[local_tool, runtime_mcp])
|
||||
|
||||
|
||||
async def test_agent_run_raises_on_duplicate_agent_mcp_names(chat_client_base: Any) -> None:
|
||||
agent = Agent(
|
||||
client=chat_client_base,
|
||||
name="TestAgent",
|
||||
tools=[
|
||||
_ConnectedMCPTool(name="docs-mcp", function_names=["search"]),
|
||||
_ConnectedMCPTool(name="github-mcp", function_names=["search"]),
|
||||
],
|
||||
)
|
||||
|
||||
with raises(ValueError, match="tool_name_prefix"):
|
||||
await agent.run("hello")
|
||||
|
||||
|
||||
async def test_agent_run_accepts_prefixed_mcp_tools(chat_client_base: Any) -> None:
|
||||
captured_options: list[dict[str, Any]] = []
|
||||
|
||||
original_inner = chat_client_base._inner_get_response
|
||||
|
||||
async def capturing_inner(
|
||||
*, messages: MutableSequence[Message], options: dict[str, Any], **kwargs: Any
|
||||
) -> ChatResponse:
|
||||
captured_options.append(dict(options))
|
||||
return await original_inner(messages=messages, options=options, **kwargs)
|
||||
|
||||
chat_client_base._inner_get_response = capturing_inner
|
||||
|
||||
local_tool = FunctionTool(func=lambda: "local", name="search", description="Local search tool")
|
||||
agent = Agent(
|
||||
client=chat_client_base,
|
||||
name="TestAgent",
|
||||
tools=[_ConnectedMCPTool(name="docs-mcp", function_names=["search"], tool_name_prefix="docs")],
|
||||
)
|
||||
|
||||
await agent.run("hello", tools=[local_tool])
|
||||
|
||||
tool_names = [tool.name for tool in captured_options[0]["tools"]]
|
||||
assert tool_names == ["search", "docs_search"]
|
||||
|
||||
|
||||
async def test_agent_tool_receives_session_in_kwargs(chat_client_base: Any) -> None:
|
||||
"""Verify tool execution receives 'session' inside **kwargs when function is called by client."""
|
||||
"""Verify legacy **kwargs tools receive the session when agent.run() is called with one."""
|
||||
|
||||
captured: dict[str, Any] = {}
|
||||
|
||||
@@ -1044,7 +1192,6 @@ async def test_agent_tool_receives_session_in_kwargs(chat_client_base: Any) -> N
|
||||
captured["has_state"] = session.state is not None if isinstance(session, AgentSession) else False
|
||||
return f"echo: {text}"
|
||||
|
||||
# Make the base client emit a function call for our tool
|
||||
chat_client_base.run_responses = [
|
||||
ChatResponse(
|
||||
messages=Message(
|
||||
@@ -1064,17 +1211,52 @@ async def test_agent_tool_receives_session_in_kwargs(chat_client_base: Any) -> N
|
||||
agent = Agent(client=chat_client_base, tools=[echo_session_info])
|
||||
session = agent.create_session()
|
||||
|
||||
result = await agent.run(
|
||||
"hello",
|
||||
session=session,
|
||||
options={"additional_function_arguments": {"session": session}},
|
||||
)
|
||||
result = await agent.run("hello", session=session)
|
||||
|
||||
assert result.text == "done"
|
||||
assert captured.get("has_session") is True
|
||||
assert captured.get("has_state") is True
|
||||
|
||||
|
||||
async def test_agent_tool_receives_explicit_session_via_function_invocation_context_kwargs(
|
||||
chat_client_base: Any,
|
||||
) -> None:
|
||||
"""Verify ctx-based tools receive the session via FunctionInvocationContext.session."""
|
||||
|
||||
captured: dict[str, Any] = {}
|
||||
|
||||
@tool(name="capture_session_context", approval_mode="never_require")
|
||||
def capture_session_context(text: str, ctx: FunctionInvocationContext) -> str:
|
||||
captured["session"] = ctx.session
|
||||
captured["has_state"] = ctx.session.state is not None if isinstance(ctx.session, AgentSession) else False
|
||||
return f"echo: {text}"
|
||||
|
||||
chat_client_base.run_responses = [
|
||||
ChatResponse(
|
||||
messages=Message(
|
||||
role="assistant",
|
||||
contents=[
|
||||
Content.from_function_call(
|
||||
call_id="1",
|
||||
name="capture_session_context",
|
||||
arguments='{"text": "hello"}',
|
||||
)
|
||||
],
|
||||
)
|
||||
),
|
||||
ChatResponse(messages=Message(role="assistant", text="done")),
|
||||
]
|
||||
|
||||
agent = Agent(client=chat_client_base, tools=[capture_session_context])
|
||||
session = agent.create_session()
|
||||
|
||||
result = await agent.run("hello", session=session)
|
||||
|
||||
assert result.text == "done"
|
||||
assert captured["session"] is session
|
||||
assert captured["has_state"] is True
|
||||
|
||||
|
||||
async def test_chat_agent_tool_choice_run_level_overrides_agent_level(chat_client_base: Any, tool_tool: Any) -> None:
|
||||
"""Verify that tool_choice passed to run() overrides agent-level tool_choice."""
|
||||
|
||||
@@ -1291,7 +1473,7 @@ def test_merge_options_none_values_ignored():
|
||||
|
||||
|
||||
def test_merge_options_tools_combined():
|
||||
"""Test _merge_options combines tool lists without duplicates."""
|
||||
"""Test _merge_options raises when distinct tools share the same name."""
|
||||
|
||||
class MockTool:
|
||||
def __init__(self, name):
|
||||
@@ -1304,13 +1486,8 @@ def test_merge_options_tools_combined():
|
||||
base = {"tools": [tool1]}
|
||||
override = {"tools": [tool2, tool3]}
|
||||
|
||||
result = _merge_options(base, override)
|
||||
|
||||
# Should have tool1 and tool2, but not duplicate tool3
|
||||
assert len(result["tools"]) == 2
|
||||
tool_names = [t.name for t in result["tools"]]
|
||||
assert "tool1" in tool_names
|
||||
assert "tool2" in tool_names
|
||||
with raises(ValueError, match="Duplicate tool name 'tool1'"):
|
||||
_merge_options(base, override)
|
||||
|
||||
|
||||
def test_merge_options_dict_tools_combined():
|
||||
@@ -1335,7 +1512,7 @@ def test_merge_options_dict_tools_combined():
|
||||
|
||||
|
||||
def test_merge_options_dict_tools_deduplicates():
|
||||
"""Test _merge_options deduplicates dict-defined tools by function name."""
|
||||
"""Test _merge_options raises on duplicate dict-defined tool names."""
|
||||
base = {
|
||||
"tools": [
|
||||
{"type": "function", "function": {"name": "tool_a"}},
|
||||
@@ -1348,12 +1525,8 @@ def test_merge_options_dict_tools_deduplicates():
|
||||
]
|
||||
}
|
||||
|
||||
result = _merge_options(base, override)
|
||||
|
||||
assert len(result["tools"]) == 2
|
||||
names = [_get_tool_name(t) for t in result["tools"]]
|
||||
assert names.count("tool_a") == 1
|
||||
assert "tool_b" in names
|
||||
with raises(ValueError, match="Duplicate tool name 'tool_a'"):
|
||||
_merge_options(base, override)
|
||||
|
||||
|
||||
def test_merge_options_mixed_tools_combined():
|
||||
@@ -1379,7 +1552,7 @@ def test_merge_options_mixed_tools_combined():
|
||||
|
||||
|
||||
def test_merge_options_mixed_tools_deduplicates():
|
||||
"""Test _merge_options deduplicates when a dict tool and object tool share the same name."""
|
||||
"""Test _merge_options raises when a dict tool and object tool share the same name."""
|
||||
|
||||
class MockTool:
|
||||
def __init__(self, name):
|
||||
@@ -1392,10 +1565,8 @@ def test_merge_options_mixed_tools_deduplicates():
|
||||
]
|
||||
}
|
||||
|
||||
result = _merge_options(base, override)
|
||||
|
||||
assert len(result["tools"]) == 1
|
||||
assert _get_tool_name(result["tools"][0]) == "tool_a"
|
||||
with raises(ValueError, match="Duplicate tool name 'tool_a'"):
|
||||
_merge_options(base, override)
|
||||
|
||||
|
||||
def test_merge_options_nameless_tools_not_deduplicated():
|
||||
@@ -1417,6 +1588,20 @@ def test_merge_options_nameless_tools_not_deduplicated():
|
||||
assert len(result["tools"]) == 2
|
||||
|
||||
|
||||
def test_merge_options_same_tool_object_kept_once():
|
||||
"""Test _merge_options silently keeps a repeated reference to the same tool object once."""
|
||||
|
||||
class MockTool:
|
||||
def __init__(self, name):
|
||||
self.name = name
|
||||
|
||||
tool_a = MockTool("tool_a")
|
||||
|
||||
result = _merge_options({"tools": [tool_a]}, {"tools": [tool_a]})
|
||||
|
||||
assert result["tools"] == [tool_a]
|
||||
|
||||
|
||||
def test_get_tool_name_dict_no_function_key():
|
||||
"""_get_tool_name returns None for a dict without a 'function' key."""
|
||||
assert _get_tool_name({"type": "function"}) is None
|
||||
@@ -1758,4 +1943,26 @@ async def test_stores_by_default_with_store_false_in_default_options_injects_inm
|
||||
assert any(isinstance(p, InMemoryHistoryProvider) for p in agent.context_providers)
|
||||
|
||||
|
||||
# endregion
|
||||
# region as_tool user_input_request propagation
|
||||
|
||||
|
||||
async def test_as_tool_raises_on_user_input_request(client: SupportsChatGetResponse) -> None:
|
||||
"""Test that as_tool raises when the wrapped sub-agent requests user input."""
|
||||
from agent_framework.exceptions import UserInputRequiredException
|
||||
|
||||
consent_content = Content.from_oauth_consent_request(
|
||||
consent_link="https://login.microsoftonline.com/consent",
|
||||
)
|
||||
client.streaming_responses = [ # type: ignore[attr-defined]
|
||||
[ChatResponseUpdate(contents=[consent_content], role="assistant")],
|
||||
]
|
||||
|
||||
agent = Agent(client=client, name="OAuthAgent", description="Agent requiring consent")
|
||||
agent_tool = agent.as_tool()
|
||||
|
||||
with raises(UserInputRequiredException) as exc_info:
|
||||
await agent_tool.invoke(arguments={"task": "Do something"})
|
||||
|
||||
assert len(exc_info.value.contents) == 1
|
||||
assert exc_info.value.contents[0].type == "oauth_consent_request"
|
||||
assert exc_info.value.contents[0].consent_link == "https://login.microsoftonline.com/consent"
|
||||
|
||||
@@ -6,7 +6,7 @@ from collections.abc import Awaitable, Callable
|
||||
from typing import Any
|
||||
|
||||
from agent_framework import Agent, ChatResponse, Content, Message, agent_middleware
|
||||
from agent_framework._middleware import AgentContext
|
||||
from agent_framework._middleware import AgentContext, FunctionInvocationContext
|
||||
|
||||
from .conftest import MockChatClient
|
||||
|
||||
@@ -14,14 +14,28 @@ from .conftest import MockChatClient
|
||||
class TestAsToolKwargsPropagation:
|
||||
"""Test cases for kwargs propagation through as_tool() delegation."""
|
||||
|
||||
@staticmethod
|
||||
def _build_context(
|
||||
tool: Any,
|
||||
*,
|
||||
task: str,
|
||||
runtime_kwargs: dict[str, Any] | None = None,
|
||||
) -> FunctionInvocationContext:
|
||||
return FunctionInvocationContext(
|
||||
function=tool,
|
||||
arguments={"task": task},
|
||||
kwargs=runtime_kwargs,
|
||||
)
|
||||
|
||||
async def test_as_tool_forwards_runtime_kwargs(self, client: MockChatClient) -> None:
|
||||
"""Test that runtime kwargs are forwarded through as_tool() to sub-agent."""
|
||||
"""Test that runtime kwargs are forwarded through as_tool() to sub-agent tools."""
|
||||
captured_kwargs: dict[str, Any] = {}
|
||||
captured_function_invocation_kwargs: dict[str, Any] = {}
|
||||
|
||||
@agent_middleware
|
||||
async def capture_middleware(context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
|
||||
# Capture kwargs passed to the sub-agent
|
||||
captured_kwargs.update(context.kwargs)
|
||||
captured_function_invocation_kwargs.update(context.function_invocation_kwargs)
|
||||
await call_next()
|
||||
|
||||
# Setup mock response
|
||||
@@ -39,29 +53,31 @@ class TestAsToolKwargsPropagation:
|
||||
# Create tool from sub-agent
|
||||
tool = sub_agent.as_tool(name="delegate", arg_name="task")
|
||||
|
||||
# Directly invoke the tool with kwargs (simulating what happens during agent execution)
|
||||
# Directly invoke the tool with explicit runtime context (simulating agent execution).
|
||||
_ = await tool.invoke(
|
||||
arguments=tool.input_model(task="Test delegation"),
|
||||
api_token="secret-xyz-123",
|
||||
user_id="user-456",
|
||||
session_id="session-789",
|
||||
context=self._build_context(
|
||||
tool,
|
||||
task="Test delegation",
|
||||
runtime_kwargs={
|
||||
"api_token": "secret-xyz-123",
|
||||
"user_id": "user-456",
|
||||
"session_id": "session-789",
|
||||
},
|
||||
),
|
||||
)
|
||||
|
||||
# Verify kwargs were forwarded to sub-agent
|
||||
assert "api_token" in captured_kwargs, f"Expected 'api_token' in {captured_kwargs}"
|
||||
assert captured_kwargs["api_token"] == "secret-xyz-123"
|
||||
assert "user_id" in captured_kwargs
|
||||
assert captured_kwargs["user_id"] == "user-456"
|
||||
assert "session_id" in captured_kwargs
|
||||
assert captured_kwargs["session_id"] == "session-789"
|
||||
assert captured_kwargs == {}
|
||||
assert captured_function_invocation_kwargs["api_token"] == "secret-xyz-123"
|
||||
assert captured_function_invocation_kwargs["user_id"] == "user-456"
|
||||
assert captured_function_invocation_kwargs["session_id"] == "session-789"
|
||||
|
||||
async def test_as_tool_excludes_arg_name_from_forwarded_kwargs(self, client: MockChatClient) -> None:
|
||||
"""Test that the arg_name parameter is not forwarded as a kwarg."""
|
||||
captured_kwargs: dict[str, Any] = {}
|
||||
async def test_as_tool_forwards_context_kwargs_verbatim(self, client: MockChatClient) -> None:
|
||||
"""Test that runtime kwargs are forwarded exactly from FunctionInvocationContext.kwargs."""
|
||||
captured_function_invocation_kwargs: dict[str, Any] = {}
|
||||
|
||||
@agent_middleware
|
||||
async def capture_middleware(context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
|
||||
captured_kwargs.update(context.kwargs)
|
||||
captured_function_invocation_kwargs.update(context.function_invocation_kwargs)
|
||||
await call_next()
|
||||
|
||||
# Setup mock response
|
||||
@@ -79,25 +95,26 @@ class TestAsToolKwargsPropagation:
|
||||
|
||||
# Invoke tool with both the arg_name field and additional kwargs
|
||||
await tool.invoke(
|
||||
arguments=tool.input_model(custom_task="Test task"),
|
||||
api_token="token-123",
|
||||
custom_task="should_be_excluded", # This should be filtered out
|
||||
context=FunctionInvocationContext(
|
||||
function=tool,
|
||||
arguments={"custom_task": "Test task"},
|
||||
kwargs={
|
||||
"api_token": "token-123",
|
||||
"custom_task": "should_be_excluded",
|
||||
},
|
||||
)
|
||||
)
|
||||
|
||||
# The arg_name ("custom_task") should NOT be in the forwarded kwargs
|
||||
assert "custom_task" not in captured_kwargs
|
||||
# But other kwargs should be present
|
||||
assert "api_token" in captured_kwargs
|
||||
assert captured_kwargs["api_token"] == "token-123"
|
||||
assert captured_function_invocation_kwargs["custom_task"] == "should_be_excluded"
|
||||
assert captured_function_invocation_kwargs["api_token"] == "token-123"
|
||||
|
||||
async def test_as_tool_nested_delegation_propagates_kwargs(self, client: MockChatClient) -> None:
|
||||
"""Test that kwargs propagate through multiple levels of delegation (A → B → C)."""
|
||||
captured_kwargs_list: list[dict[str, Any]] = []
|
||||
"""Test that runtime kwargs propagate through multiple levels of delegation (A -> B -> C)."""
|
||||
captured_function_invocation_kwargs_list: list[dict[str, Any]] = []
|
||||
|
||||
@agent_middleware
|
||||
async def capture_middleware(context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
|
||||
# Capture kwargs at each level
|
||||
captured_kwargs_list.append(dict(context.kwargs))
|
||||
captured_function_invocation_kwargs_list.append(dict(context.function_invocation_kwargs))
|
||||
await call_next()
|
||||
|
||||
# Setup mock responses to trigger nested tool invocation: B calls tool C, then completes.
|
||||
@@ -140,24 +157,29 @@ class TestAsToolKwargsPropagation:
|
||||
|
||||
# Invoke tool B with kwargs - should propagate to both B and C
|
||||
await tool_b.invoke(
|
||||
arguments=tool_b.input_model(task="Test cascade"),
|
||||
trace_id="trace-abc-123",
|
||||
tenant_id="tenant-xyz",
|
||||
options={"additional_function_arguments": {"trace_id": "trace-abc-123", "tenant_id": "tenant-xyz"}},
|
||||
context=self._build_context(
|
||||
tool_b,
|
||||
task="Test cascade",
|
||||
runtime_kwargs={
|
||||
"trace_id": "trace-abc-123",
|
||||
"tenant_id": "tenant-xyz",
|
||||
},
|
||||
),
|
||||
)
|
||||
|
||||
# Verify kwargs were forwarded to the first agent invocation.
|
||||
assert len(captured_kwargs_list) >= 1
|
||||
assert captured_kwargs_list[0].get("trace_id") == "trace-abc-123"
|
||||
assert captured_kwargs_list[0].get("tenant_id") == "tenant-xyz"
|
||||
assert len(captured_function_invocation_kwargs_list) >= 1
|
||||
assert captured_function_invocation_kwargs_list[0].get("trace_id") == "trace-abc-123"
|
||||
assert captured_function_invocation_kwargs_list[0].get("tenant_id") == "tenant-xyz"
|
||||
|
||||
async def test_as_tool_streaming_mode_forwards_kwargs(self, client: MockChatClient) -> None:
|
||||
"""Test that kwargs are forwarded in streaming mode."""
|
||||
"""Test that runtime kwargs are forwarded in streaming mode."""
|
||||
captured_kwargs: dict[str, Any] = {}
|
||||
captured_function_invocation_kwargs: dict[str, Any] = {}
|
||||
|
||||
@agent_middleware
|
||||
async def capture_middleware(context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
|
||||
captured_kwargs.update(context.kwargs)
|
||||
captured_function_invocation_kwargs.update(context.function_invocation_kwargs)
|
||||
await call_next()
|
||||
|
||||
# Setup mock streaming responses
|
||||
@@ -182,13 +204,15 @@ class TestAsToolKwargsPropagation:
|
||||
|
||||
# Invoke tool with kwargs while streaming callback is active
|
||||
await tool.invoke(
|
||||
arguments=tool.input_model(task="Test streaming"),
|
||||
api_key="streaming-key-999",
|
||||
context=self._build_context(
|
||||
tool,
|
||||
task="Test streaming",
|
||||
runtime_kwargs={"api_key": "streaming-key-999"},
|
||||
),
|
||||
)
|
||||
|
||||
# Verify kwargs were forwarded even in streaming mode
|
||||
assert "api_key" in captured_kwargs
|
||||
assert captured_kwargs["api_key"] == "streaming-key-999"
|
||||
assert captured_kwargs == {}
|
||||
assert captured_function_invocation_kwargs["api_key"] == "streaming-key-999"
|
||||
assert len(captured_updates) == 1
|
||||
|
||||
async def test_as_tool_empty_kwargs_still_works(self, client: MockChatClient) -> None:
|
||||
@@ -206,18 +230,20 @@ class TestAsToolKwargsPropagation:
|
||||
tool = sub_agent.as_tool()
|
||||
|
||||
# Invoke without any extra kwargs - should work without errors
|
||||
result = await tool.invoke(arguments=tool.input_model(task="Simple task"))
|
||||
result = await tool.invoke(arguments={"task": "Simple task"})
|
||||
|
||||
# Verify tool executed successfully
|
||||
assert result is not None
|
||||
|
||||
async def test_as_tool_kwargs_with_chat_options(self, client: MockChatClient) -> None:
|
||||
"""Test that kwargs including chat_options are properly forwarded."""
|
||||
"""Test that runtime kwargs are forwarded only via function_invocation_kwargs."""
|
||||
captured_kwargs: dict[str, Any] = {}
|
||||
captured_function_invocation_kwargs: dict[str, Any] = {}
|
||||
|
||||
@agent_middleware
|
||||
async def capture_middleware(context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
|
||||
captured_kwargs.update(context.kwargs)
|
||||
captured_function_invocation_kwargs.update(context.function_invocation_kwargs)
|
||||
await call_next()
|
||||
|
||||
# Setup mock response
|
||||
@@ -235,24 +261,26 @@ class TestAsToolKwargsPropagation:
|
||||
|
||||
# Invoke with various kwargs
|
||||
await tool.invoke(
|
||||
arguments=tool.input_model(task="Test with options"),
|
||||
temperature=0.8,
|
||||
max_tokens=500,
|
||||
custom_param="custom_value",
|
||||
context=self._build_context(
|
||||
tool,
|
||||
task="Test with options",
|
||||
runtime_kwargs={
|
||||
"temperature": 0.8,
|
||||
"max_tokens": 500,
|
||||
"custom_param": "custom_value",
|
||||
},
|
||||
),
|
||||
)
|
||||
|
||||
# Verify all kwargs were forwarded
|
||||
assert "temperature" in captured_kwargs
|
||||
assert captured_kwargs["temperature"] == 0.8
|
||||
assert "max_tokens" in captured_kwargs
|
||||
assert captured_kwargs["max_tokens"] == 500
|
||||
assert "custom_param" in captured_kwargs
|
||||
assert captured_kwargs["custom_param"] == "custom_value"
|
||||
assert captured_kwargs == {}
|
||||
assert captured_function_invocation_kwargs["temperature"] == 0.8
|
||||
assert captured_function_invocation_kwargs["max_tokens"] == 500
|
||||
assert captured_function_invocation_kwargs["custom_param"] == "custom_value"
|
||||
|
||||
async def test_as_tool_kwargs_isolated_per_invocation(self, client: MockChatClient) -> None:
|
||||
"""Test that kwargs are isolated per invocation and don't leak between calls."""
|
||||
first_call_kwargs: dict[str, Any] = {}
|
||||
second_call_kwargs: dict[str, Any] = {}
|
||||
"""Test that runtime kwargs are isolated per invocation and don't leak between calls."""
|
||||
first_call_function_invocation_kwargs: dict[str, Any] = {}
|
||||
second_call_function_invocation_kwargs: dict[str, Any] = {}
|
||||
call_count = 0
|
||||
|
||||
@agent_middleware
|
||||
@@ -260,9 +288,9 @@ class TestAsToolKwargsPropagation:
|
||||
nonlocal call_count
|
||||
call_count += 1
|
||||
if call_count == 1:
|
||||
first_call_kwargs.update(context.kwargs)
|
||||
first_call_function_invocation_kwargs.update(context.function_invocation_kwargs)
|
||||
elif call_count == 2:
|
||||
second_call_kwargs.update(context.kwargs)
|
||||
second_call_function_invocation_kwargs.update(context.function_invocation_kwargs)
|
||||
await call_next()
|
||||
|
||||
# Setup mock responses for both calls
|
||||
@@ -281,33 +309,35 @@ class TestAsToolKwargsPropagation:
|
||||
|
||||
# First call with specific kwargs
|
||||
await tool.invoke(
|
||||
arguments=tool.input_model(task="First task"),
|
||||
session_id="session-1",
|
||||
api_token="token-1",
|
||||
context=self._build_context(
|
||||
tool,
|
||||
task="First task",
|
||||
runtime_kwargs={"session_id": "session-1", "api_token": "token-1"},
|
||||
),
|
||||
)
|
||||
|
||||
# Second call with different kwargs
|
||||
await tool.invoke(
|
||||
arguments=tool.input_model(task="Second task"),
|
||||
session_id="session-2",
|
||||
api_token="token-2",
|
||||
context=self._build_context(
|
||||
tool,
|
||||
task="Second task",
|
||||
runtime_kwargs={"session_id": "session-2", "api_token": "token-2"},
|
||||
),
|
||||
)
|
||||
|
||||
# Verify first call had its own kwargs
|
||||
assert first_call_kwargs.get("session_id") == "session-1"
|
||||
assert first_call_kwargs.get("api_token") == "token-1"
|
||||
assert first_call_function_invocation_kwargs.get("session_id") == "session-1"
|
||||
assert first_call_function_invocation_kwargs.get("api_token") == "token-1"
|
||||
|
||||
# Verify second call had its own kwargs (not leaked from first)
|
||||
assert second_call_kwargs.get("session_id") == "session-2"
|
||||
assert second_call_kwargs.get("api_token") == "token-2"
|
||||
assert second_call_function_invocation_kwargs.get("session_id") == "session-2"
|
||||
assert second_call_function_invocation_kwargs.get("api_token") == "token-2"
|
||||
|
||||
async def test_as_tool_excludes_conversation_id_from_forwarded_kwargs(self, client: MockChatClient) -> None:
|
||||
"""Test that conversation_id is not forwarded to sub-agent."""
|
||||
captured_kwargs: dict[str, Any] = {}
|
||||
async def test_as_tool_forwards_conversation_id_from_context_kwargs(self, client: MockChatClient) -> None:
|
||||
"""Test that conversation_id is forwarded when explicitly present in runtime context kwargs."""
|
||||
captured_function_invocation_kwargs: dict[str, Any] = {}
|
||||
|
||||
@agent_middleware
|
||||
async def capture_middleware(context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
|
||||
captured_kwargs.update(context.kwargs)
|
||||
captured_function_invocation_kwargs.update(context.function_invocation_kwargs)
|
||||
await call_next()
|
||||
|
||||
# Setup mock response
|
||||
@@ -325,17 +355,17 @@ class TestAsToolKwargsPropagation:
|
||||
|
||||
# Invoke tool with conversation_id in kwargs (simulating parent's conversation state)
|
||||
await tool.invoke(
|
||||
arguments=tool.input_model(task="Test delegation"),
|
||||
conversation_id="conv-parent-456",
|
||||
api_token="secret-xyz-123",
|
||||
user_id="user-456",
|
||||
context=self._build_context(
|
||||
tool,
|
||||
task="Test delegation",
|
||||
runtime_kwargs={
|
||||
"conversation_id": "conv-parent-456",
|
||||
"api_token": "secret-xyz-123",
|
||||
"user_id": "user-456",
|
||||
},
|
||||
),
|
||||
)
|
||||
|
||||
# Verify conversation_id was NOT forwarded to sub-agent
|
||||
assert "conversation_id" not in captured_kwargs, (
|
||||
f"conversation_id should not be forwarded, but got: {captured_kwargs}"
|
||||
)
|
||||
|
||||
# Verify other kwargs were still forwarded
|
||||
assert captured_kwargs.get("api_token") == "secret-xyz-123"
|
||||
assert captured_kwargs.get("user_id") == "user-456"
|
||||
assert captured_function_invocation_kwargs.get("conversation_id") == "conv-parent-456"
|
||||
assert captured_function_invocation_kwargs.get("api_token") == "secret-xyz-123"
|
||||
assert captured_function_invocation_kwargs.get("user_id") == "user-456"
|
||||
|
||||
@@ -1,9 +1,12 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
|
||||
import inspect
|
||||
from typing import Any
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
|
||||
from agent_framework import (
|
||||
GROUP_ANNOTATION_KEY,
|
||||
GROUP_TOKEN_COUNT_KEY,
|
||||
@@ -50,6 +53,60 @@ def test_base_client(chat_client_base: SupportsChatGetResponse):
|
||||
assert isinstance(chat_client_base, SupportsChatGetResponse)
|
||||
|
||||
|
||||
def test_base_client_warns_for_direct_additional_properties(chat_client_base: SupportsChatGetResponse) -> None:
|
||||
with pytest.warns(DeprecationWarning, match="additional_properties"):
|
||||
client = type(chat_client_base)(legacy_key="legacy-value")
|
||||
|
||||
assert client.additional_properties["legacy_key"] == "legacy-value"
|
||||
|
||||
|
||||
def test_base_client_as_agent_uses_explicit_additional_properties(chat_client_base: SupportsChatGetResponse) -> None:
|
||||
agent = chat_client_base.as_agent(additional_properties={"team": "core"})
|
||||
|
||||
assert agent.additional_properties == {"team": "core"}
|
||||
|
||||
|
||||
def test_openai_chat_client_get_response_docstring_surfaces_layered_runtime_docs() -> None:
|
||||
from agent_framework.openai import OpenAIChatClient
|
||||
|
||||
docstring = inspect.getdoc(OpenAIChatClient.get_response)
|
||||
|
||||
assert docstring is not None
|
||||
assert "Get a response from a chat client." in docstring
|
||||
assert "function_invocation_kwargs" in docstring
|
||||
assert "function_middleware: Optional per-call function middleware." in docstring
|
||||
assert "middleware: Optional per-call chat and function middleware." in docstring
|
||||
|
||||
|
||||
def test_openai_chat_client_get_response_is_defined_on_openai_class() -> None:
|
||||
from agent_framework.openai import OpenAIChatClient
|
||||
|
||||
signature = inspect.signature(OpenAIChatClient.get_response)
|
||||
|
||||
assert OpenAIChatClient.get_response.__qualname__ == "OpenAIChatClient.get_response"
|
||||
assert "function_middleware" in signature.parameters
|
||||
assert "middleware" in signature.parameters
|
||||
|
||||
|
||||
async def test_base_client_get_response_uses_explicit_client_kwargs(chat_client_base: SupportsChatGetResponse) -> None:
|
||||
async def fake_inner_get_response(**kwargs):
|
||||
assert kwargs["trace_id"] == "trace-123"
|
||||
assert "function_invocation_kwargs" not in kwargs
|
||||
return ChatResponse(messages=[Message(role="assistant", text="ok")])
|
||||
|
||||
with patch.object(
|
||||
chat_client_base,
|
||||
"_inner_get_response",
|
||||
side_effect=fake_inner_get_response,
|
||||
) as mock_inner_get_response:
|
||||
await chat_client_base.get_response(
|
||||
[Message(role="user", text="hello")],
|
||||
function_invocation_kwargs={"tool_request_id": "tool-123"},
|
||||
client_kwargs={"trace_id": "trace-123"},
|
||||
)
|
||||
mock_inner_get_response.assert_called_once()
|
||||
|
||||
|
||||
async def test_base_client_get_response(chat_client_base: SupportsChatGetResponse):
|
||||
response = await chat_client_base.get_response([Message(role="user", text="Hello")])
|
||||
assert response.messages[0].role == "assistant"
|
||||
|
||||
@@ -0,0 +1,175 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
from agent_framework._docstrings import apply_layered_docstring, build_layered_docstring
|
||||
|
||||
# -- Helpers: stub functions with various docstring shapes --
|
||||
|
||||
|
||||
def _source_with_full_docstring(x: int) -> int:
|
||||
"""Do something useful.
|
||||
|
||||
Args:
|
||||
x: The input value.
|
||||
|
||||
Keyword Args:
|
||||
timeout: Max seconds to wait.
|
||||
|
||||
Returns:
|
||||
The computed result.
|
||||
"""
|
||||
return x
|
||||
|
||||
|
||||
def _source_with_args_only(x: int) -> int:
|
||||
"""Do something useful.
|
||||
|
||||
Args:
|
||||
x: The input value.
|
||||
|
||||
Returns:
|
||||
The computed result.
|
||||
"""
|
||||
return x
|
||||
|
||||
|
||||
def _source_no_sections() -> None:
|
||||
"""A plain summary with no Google-style sections."""
|
||||
|
||||
|
||||
def _source_no_docstring() -> None:
|
||||
pass
|
||||
|
||||
|
||||
def _target_stub() -> None:
|
||||
pass
|
||||
|
||||
|
||||
# -- build_layered_docstring tests --
|
||||
|
||||
|
||||
def test_build_returns_none_when_source_has_no_docstring() -> None:
|
||||
result = build_layered_docstring(_source_no_docstring)
|
||||
assert result is None
|
||||
|
||||
|
||||
def test_build_returns_original_when_no_extra_kwargs() -> None:
|
||||
result = build_layered_docstring(_source_with_full_docstring)
|
||||
assert result is not None
|
||||
assert "Do something useful." in result
|
||||
assert "Keyword Args:" in result
|
||||
|
||||
|
||||
def test_build_returns_original_when_extra_kwargs_empty() -> None:
|
||||
result = build_layered_docstring(_source_with_full_docstring, extra_keyword_args={})
|
||||
assert result is not None
|
||||
assert result == build_layered_docstring(_source_with_full_docstring)
|
||||
|
||||
|
||||
def test_build_appends_to_existing_keyword_args_section() -> None:
|
||||
result = build_layered_docstring(
|
||||
_source_with_full_docstring,
|
||||
extra_keyword_args={"retries": "Number of retries."},
|
||||
)
|
||||
assert result is not None
|
||||
assert "timeout: Max seconds to wait." in result
|
||||
assert "retries: Number of retries." in result
|
||||
# Both should be under Keyword Args
|
||||
lines = result.splitlines()
|
||||
kw_index = next(i for i, line in enumerate(lines) if line == "Keyword Args:")
|
||||
ret_index = next(i for i, line in enumerate(lines) if line == "Returns:")
|
||||
retries_index = next(i for i, line in enumerate(lines) if "retries:" in line)
|
||||
assert kw_index < retries_index < ret_index
|
||||
|
||||
|
||||
def test_build_inserts_keyword_args_after_args_section() -> None:
|
||||
result = build_layered_docstring(
|
||||
_source_with_args_only,
|
||||
extra_keyword_args={"verbose": "Enable verbose output."},
|
||||
)
|
||||
assert result is not None
|
||||
assert "Keyword Args:" in result
|
||||
assert "verbose: Enable verbose output." in result
|
||||
lines = result.splitlines()
|
||||
args_index = next(i for i, line in enumerate(lines) if line == "Args:")
|
||||
kw_index = next(i for i, line in enumerate(lines) if line == "Keyword Args:")
|
||||
ret_index = next(i for i, line in enumerate(lines) if line == "Returns:")
|
||||
assert args_index < kw_index < ret_index
|
||||
|
||||
|
||||
def test_build_inserts_keyword_args_in_docstring_with_no_sections() -> None:
|
||||
result = build_layered_docstring(
|
||||
_source_no_sections,
|
||||
extra_keyword_args={"debug": "Enable debug mode."},
|
||||
)
|
||||
assert result is not None
|
||||
assert "A plain summary" in result
|
||||
assert "Keyword Args:" in result
|
||||
assert "debug: Enable debug mode." in result
|
||||
|
||||
|
||||
def test_build_handles_multiline_descriptions() -> None:
|
||||
result = build_layered_docstring(
|
||||
_source_with_args_only,
|
||||
extra_keyword_args={
|
||||
"config": "The configuration object.\nMust be a valid mapping.\nDefaults to empty.",
|
||||
},
|
||||
)
|
||||
assert result is not None
|
||||
lines = result.splitlines()
|
||||
config_line = next(line for line in lines if "config:" in line)
|
||||
assert "The configuration object." in config_line
|
||||
# Continuation lines should be indented
|
||||
config_idx = lines.index(config_line)
|
||||
assert "Must be a valid mapping." in lines[config_idx + 1]
|
||||
assert "Defaults to empty." in lines[config_idx + 2]
|
||||
|
||||
|
||||
def test_build_preserves_multiple_extra_kwargs_order() -> None:
|
||||
result = build_layered_docstring(
|
||||
_source_with_args_only,
|
||||
extra_keyword_args={
|
||||
"alpha": "First.",
|
||||
"beta": "Second.",
|
||||
"gamma": "Third.",
|
||||
},
|
||||
)
|
||||
assert result is not None
|
||||
lines = result.splitlines()
|
||||
alpha_idx = next(i for i, line in enumerate(lines) if "alpha:" in line)
|
||||
beta_idx = next(i for i, line in enumerate(lines) if "beta:" in line)
|
||||
gamma_idx = next(i for i, line in enumerate(lines) if "gamma:" in line)
|
||||
assert alpha_idx < beta_idx < gamma_idx
|
||||
|
||||
|
||||
# -- apply_layered_docstring tests --
|
||||
|
||||
|
||||
def test_apply_sets_docstring_on_target() -> None:
|
||||
def target() -> None:
|
||||
pass
|
||||
|
||||
apply_layered_docstring(target, _source_with_full_docstring)
|
||||
assert target.__doc__ is not None
|
||||
assert "Do something useful." in target.__doc__
|
||||
|
||||
|
||||
def test_apply_with_extra_kwargs() -> None:
|
||||
def target() -> None:
|
||||
pass
|
||||
|
||||
apply_layered_docstring(
|
||||
target,
|
||||
_source_with_args_only,
|
||||
extra_keyword_args={"flag": "A boolean flag."},
|
||||
)
|
||||
assert target.__doc__ is not None
|
||||
assert "flag: A boolean flag." in target.__doc__
|
||||
assert "Keyword Args:" in target.__doc__
|
||||
|
||||
|
||||
def test_apply_sets_none_when_source_has_no_docstring() -> None:
|
||||
def target() -> None:
|
||||
"""Original."""
|
||||
|
||||
apply_layered_docstring(target, _source_no_docstring)
|
||||
assert target.__doc__ is None
|
||||
@@ -4,6 +4,8 @@ from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
import pytest
|
||||
|
||||
from agent_framework import (
|
||||
BaseEmbeddingClient,
|
||||
Embedding,
|
||||
@@ -63,6 +65,11 @@ def test_base_additional_properties_custom() -> None:
|
||||
assert client.additional_properties == {"key": "value"}
|
||||
|
||||
|
||||
def test_base_embedding_client_rejects_unknown_kwargs() -> None:
|
||||
with pytest.raises(TypeError):
|
||||
MockEmbeddingClient(legacy_key="value") # type: ignore[call-arg]
|
||||
|
||||
|
||||
# --- SupportsGetEmbeddings protocol tests ---
|
||||
|
||||
|
||||
|
||||
@@ -3651,3 +3651,131 @@ class TestUpdateConversationId:
|
||||
|
||||
|
||||
# endregion
|
||||
async def test_user_input_request_propagates_through_as_tool(chat_client_base: SupportsChatGetResponse):
|
||||
"""Test that user_input_request content from a sub-agent wrapped as a tool propagates to the parent response."""
|
||||
from agent_framework.exceptions import UserInputRequiredException
|
||||
|
||||
@tool(name="delegate_agent", approval_mode="never_require")
|
||||
def delegate_tool(task: str) -> str:
|
||||
del task
|
||||
raise UserInputRequiredException(
|
||||
contents=[
|
||||
Content.from_oauth_consent_request(
|
||||
consent_link="https://login.microsoftonline.com/consent",
|
||||
)
|
||||
]
|
||||
)
|
||||
|
||||
chat_client_base.run_responses = [
|
||||
ChatResponse(
|
||||
messages=Message(
|
||||
role="assistant",
|
||||
contents=[
|
||||
Content.from_function_call(call_id="1", name="delegate_agent", arguments='{"task": "do it"}'),
|
||||
],
|
||||
)
|
||||
)
|
||||
]
|
||||
|
||||
response = await chat_client_base.get_response(
|
||||
[Message(role="user", text="delegate this")],
|
||||
options={"tool_choice": "auto", "tools": [delegate_tool]},
|
||||
)
|
||||
|
||||
user_requests = [
|
||||
content
|
||||
for msg in response.messages
|
||||
for content in msg.contents
|
||||
if isinstance(content, Content) and content.user_input_request
|
||||
]
|
||||
assert len(user_requests) == 1
|
||||
assert user_requests[0].type == "oauth_consent_request"
|
||||
assert user_requests[0].consent_link == "https://login.microsoftonline.com/consent"
|
||||
assert user_requests[0].user_input_request is True
|
||||
|
||||
|
||||
async def test_user_input_request_multiple_contents_propagate(chat_client_base: SupportsChatGetResponse):
|
||||
"""Test that multiple user_input_request items in a single exception all propagate to the parent response."""
|
||||
from agent_framework.exceptions import UserInputRequiredException
|
||||
|
||||
@tool(name="multi_request_tool", approval_mode="never_require")
|
||||
def multi_request(task: str) -> str:
|
||||
del task
|
||||
raise UserInputRequiredException(
|
||||
contents=[
|
||||
Content.from_oauth_consent_request(
|
||||
consent_link="https://example.com/consent1",
|
||||
),
|
||||
Content.from_oauth_consent_request(
|
||||
consent_link="https://example.com/consent2",
|
||||
),
|
||||
Content.from_oauth_consent_request(
|
||||
consent_link="https://example.com/consent3",
|
||||
),
|
||||
]
|
||||
)
|
||||
|
||||
chat_client_base.run_responses = [
|
||||
ChatResponse(
|
||||
messages=Message(
|
||||
role="assistant",
|
||||
contents=[
|
||||
Content.from_function_call(call_id="1", name="multi_request_tool", arguments='{"task": "do it"}'),
|
||||
],
|
||||
)
|
||||
)
|
||||
]
|
||||
|
||||
response = await chat_client_base.get_response(
|
||||
[Message(role="user", text="do something")],
|
||||
options={"tool_choice": "auto", "tools": [multi_request]},
|
||||
)
|
||||
|
||||
user_requests = [
|
||||
content
|
||||
for msg in response.messages
|
||||
for content in msg.contents
|
||||
if isinstance(content, Content) and content.user_input_request
|
||||
]
|
||||
assert len(user_requests) == 3
|
||||
consent_links = {r.consent_link for r in user_requests}
|
||||
assert consent_links == {
|
||||
"https://example.com/consent1",
|
||||
"https://example.com/consent2",
|
||||
"https://example.com/consent3",
|
||||
}
|
||||
|
||||
|
||||
async def test_user_input_request_empty_contents_returns_fallback(chat_client_base: SupportsChatGetResponse):
|
||||
"""Test that UserInputRequiredException with empty contents produces a fallback function_result."""
|
||||
from agent_framework.exceptions import UserInputRequiredException
|
||||
|
||||
@tool(name="empty_request_tool", approval_mode="never_require")
|
||||
def empty_request(task: str) -> str:
|
||||
del task
|
||||
raise UserInputRequiredException(contents=[])
|
||||
|
||||
chat_client_base.run_responses = [
|
||||
ChatResponse(
|
||||
messages=Message(
|
||||
role="assistant",
|
||||
contents=[
|
||||
Content.from_function_call(call_id="1", name="empty_request_tool", arguments='{"task": "do it"}'),
|
||||
],
|
||||
)
|
||||
),
|
||||
ChatResponse(messages=Message(role="assistant", text="handled")),
|
||||
]
|
||||
|
||||
response = await chat_client_base.get_response(
|
||||
[Message(role="user", text="do something")],
|
||||
options={"tool_choice": "auto", "tools": [empty_request]},
|
||||
)
|
||||
|
||||
# With empty contents, the handler returns a function_result with an error message
|
||||
# and the loop continues to the next chat response.
|
||||
function_results = [
|
||||
content for msg in response.messages for content in msg.contents if content.type == "function_result"
|
||||
]
|
||||
assert len(function_results) >= 1
|
||||
assert any("user input" in (fr.result or "").lower() for fr in function_results)
|
||||
|
||||
@@ -6,11 +6,13 @@ from collections.abc import AsyncIterable, Awaitable, MutableSequence, Sequence
|
||||
from typing import Any
|
||||
|
||||
from agent_framework import (
|
||||
Agent,
|
||||
BaseChatClient,
|
||||
ChatMiddlewareLayer,
|
||||
ChatResponse,
|
||||
ChatResponseUpdate,
|
||||
Content,
|
||||
FunctionInvocationContext,
|
||||
FunctionInvocationLayer,
|
||||
Message,
|
||||
ResponseStream,
|
||||
@@ -97,6 +99,7 @@ class TestKwargsPropagationToFunctionTool:
|
||||
|
||||
async def test_kwargs_propagate_to_tool_with_kwargs(self) -> None:
|
||||
"""Test that kwargs passed to get_response() are available in @tool **kwargs."""
|
||||
# TODO(Copilot): Remove this legacy coverage once runtime ``**kwargs`` tool injection is removed.
|
||||
captured_kwargs: dict[str, Any] = {}
|
||||
|
||||
@tool(approval_mode="never_require")
|
||||
@@ -149,6 +152,7 @@ class TestKwargsPropagationToFunctionTool:
|
||||
|
||||
async def test_kwargs_not_forwarded_to_tool_without_kwargs(self) -> None:
|
||||
"""Test that kwargs are NOT forwarded to @tool that doesn't accept **kwargs."""
|
||||
# TODO(Copilot): Remove this legacy coverage once runtime ``**kwargs`` tool injection is removed.
|
||||
|
||||
@tool(approval_mode="never_require")
|
||||
def simple_tool(x: int) -> str:
|
||||
@@ -185,6 +189,7 @@ class TestKwargsPropagationToFunctionTool:
|
||||
|
||||
async def test_kwargs_isolated_between_function_calls(self) -> None:
|
||||
"""Test that kwargs are consistent across multiple function call invocations."""
|
||||
# TODO(Copilot): Remove this legacy coverage once runtime ``**kwargs`` tool injection is removed.
|
||||
invocation_kwargs: list[dict[str, Any]] = []
|
||||
|
||||
@tool(approval_mode="never_require")
|
||||
@@ -235,6 +240,7 @@ class TestKwargsPropagationToFunctionTool:
|
||||
|
||||
async def test_streaming_response_kwargs_propagation(self) -> None:
|
||||
"""Test that kwargs propagate to @tool in streaming mode."""
|
||||
# TODO(Copilot): Remove this legacy coverage once runtime ``**kwargs`` tool injection is removed.
|
||||
captured_kwargs: dict[str, Any] = {}
|
||||
|
||||
@tool(approval_mode="never_require")
|
||||
@@ -287,3 +293,59 @@ class TestKwargsPropagationToFunctionTool:
|
||||
assert "streaming_session" in captured_kwargs, f"Expected 'streaming_session' in {captured_kwargs}"
|
||||
assert captured_kwargs["streaming_session"] == "session-xyz"
|
||||
assert captured_kwargs["correlation_id"] == "corr-123"
|
||||
|
||||
async def test_agent_run_injects_function_invocation_context(self) -> None:
|
||||
"""Test that Agent.run injects FunctionInvocationContext for ctx-based tools."""
|
||||
captured_context_kwargs: dict[str, Any] = {}
|
||||
captured_client_kwargs: dict[str, Any] = {}
|
||||
captured_options: dict[str, Any] = {}
|
||||
|
||||
@tool(approval_mode="never_require")
|
||||
def capture_context_tool(x: int, ctx: FunctionInvocationContext) -> str:
|
||||
captured_context_kwargs.update(ctx.kwargs)
|
||||
return f"result: x={x}"
|
||||
|
||||
class CapturingFunctionInvokingMockClient(FunctionInvokingMockClient):
|
||||
async def _get_non_streaming_response(
|
||||
self,
|
||||
*,
|
||||
messages: MutableSequence[Message],
|
||||
options: dict[str, Any],
|
||||
**kwargs: Any,
|
||||
) -> ChatResponse:
|
||||
captured_options.update(options)
|
||||
captured_client_kwargs.update(kwargs)
|
||||
return await super()._get_non_streaming_response(messages=messages, options=options, **kwargs)
|
||||
|
||||
client = CapturingFunctionInvokingMockClient()
|
||||
client.run_responses = [
|
||||
ChatResponse(
|
||||
messages=[
|
||||
Message(
|
||||
role="assistant",
|
||||
contents=[
|
||||
Content.from_function_call(
|
||||
call_id="call_1",
|
||||
name="capture_context_tool",
|
||||
arguments='{"x": 42}',
|
||||
)
|
||||
],
|
||||
)
|
||||
]
|
||||
),
|
||||
ChatResponse(messages=[Message(role="assistant", text="Done!")]),
|
||||
]
|
||||
|
||||
agent = Agent(client=client, tools=[capture_context_tool])
|
||||
result = await agent.run(
|
||||
[Message(role="user", text="Test")],
|
||||
function_invocation_kwargs={"tool_request_id": "tool-123"},
|
||||
client_kwargs={"client_request_id": "client-456"},
|
||||
)
|
||||
|
||||
assert captured_context_kwargs["tool_request_id"] == "tool-123"
|
||||
assert "client_request_id" not in captured_context_kwargs
|
||||
assert captured_client_kwargs["client_request_id"] == "client-456"
|
||||
assert "tool_request_id" not in captured_client_kwargs
|
||||
assert "additional_function_arguments" not in captured_options
|
||||
assert result.messages[-1].text == "Done!"
|
||||
|
||||
@@ -53,6 +53,81 @@ def test_normalize_mcp_name():
|
||||
assert _normalize_mcp_name("name/with\\slashes") == "name-with-slashes"
|
||||
|
||||
|
||||
def test_mcp_transport_subclasses_accept_tool_name_prefix() -> None:
|
||||
assert MCPStdioTool(name="stdio", command="python", tool_name_prefix="stdio").tool_name_prefix == "stdio"
|
||||
assert (
|
||||
MCPStreamableHTTPTool(
|
||||
name="http",
|
||||
url="https://example.com/mcp",
|
||||
tool_name_prefix="http",
|
||||
).tool_name_prefix
|
||||
== "http"
|
||||
)
|
||||
assert (
|
||||
MCPWebsocketTool(
|
||||
name="ws",
|
||||
url="wss://example.com/mcp",
|
||||
tool_name_prefix="ws",
|
||||
).tool_name_prefix
|
||||
== "ws"
|
||||
)
|
||||
|
||||
|
||||
async def test_load_tools_with_tool_name_prefix_preserves_matching_configuration():
|
||||
"""Prefixed MCP tool names should still honor unprefixed allow/approval configuration."""
|
||||
tool = MCPTool(
|
||||
name="docs",
|
||||
tool_name_prefix="docs",
|
||||
allowed_tools=["search_docs"],
|
||||
approval_mode={"always_require_approval": ["search_docs"]},
|
||||
)
|
||||
|
||||
mock_session = AsyncMock()
|
||||
tool.session = mock_session
|
||||
tool.load_tools_flag = True
|
||||
|
||||
page = Mock()
|
||||
page.tools = [
|
||||
types.Tool(
|
||||
name="search_docs",
|
||||
description="Search docs",
|
||||
inputSchema={"type": "object", "properties": {"query": {"type": "string"}}},
|
||||
),
|
||||
]
|
||||
page.nextCursor = None
|
||||
mock_session.list_tools = AsyncMock(return_value=page)
|
||||
|
||||
await tool.load_tools()
|
||||
|
||||
assert [function.name for function in tool._functions] == ["docs_search_docs"]
|
||||
assert [function.name for function in tool.functions] == ["docs_search_docs"]
|
||||
assert tool.functions[0].approval_mode == "always_require"
|
||||
|
||||
|
||||
async def test_load_prompts_with_tool_name_prefix() -> None:
|
||||
"""Prefixed MCP prompt names should be exposed with the configured prefix."""
|
||||
tool = MCPTool(name="docs", tool_name_prefix="docs")
|
||||
|
||||
mock_session = AsyncMock()
|
||||
tool.session = mock_session
|
||||
tool.load_prompts_flag = True
|
||||
|
||||
page = Mock()
|
||||
page.prompts = [
|
||||
types.Prompt(
|
||||
name="summarize docs",
|
||||
description="Summarize docs",
|
||||
arguments=[types.PromptArgument(name="topic", description="Topic", required=True)],
|
||||
),
|
||||
]
|
||||
page.nextCursor = None
|
||||
mock_session.list_prompts = AsyncMock(return_value=page)
|
||||
|
||||
await tool.load_prompts()
|
||||
|
||||
assert [function.name for function in tool._functions] == ["docs_summarize-docs"]
|
||||
|
||||
|
||||
def test_mcp_prompt_message_to_ai_content():
|
||||
"""Test conversion from MCP prompt message to AI content."""
|
||||
mcp_message = types.PromptMessage(role="user", content=types.TextContent(type="text", text="Hello, world!"))
|
||||
@@ -120,13 +195,16 @@ def test_parse_tool_result_from_mcp_meta_not_in_string():
|
||||
|
||||
|
||||
def test_parse_tool_result_from_mcp_empty_content():
|
||||
"""Test that empty content produces list with empty text Content."""
|
||||
"""Test that empty MCP content normalizes to JSON null text content."""
|
||||
mcp_result = types.CallToolResult(content=[])
|
||||
result = _parse_tool_result_from_mcp(mcp_result)
|
||||
assert isinstance(result, list)
|
||||
assert len(result) == 1
|
||||
assert result[0].type == "text"
|
||||
assert result[0].text == ""
|
||||
assert result[0].text == "null"
|
||||
|
||||
function_result = Content.from_function_result(call_id="call_null", result=result)
|
||||
assert function_result.result == "null"
|
||||
|
||||
|
||||
def test_parse_tool_result_from_mcp_audio_content():
|
||||
@@ -2447,67 +2525,169 @@ async def test_mcp_tool_get_prompt_reconnection_on_closed_resource_error():
|
||||
assert "failed to reconnect" in str(exc_info.value).lower()
|
||||
|
||||
|
||||
async def test_mcp_tool_reconnection_handles_cross_task_cancel_scope_error():
|
||||
"""Test that reconnection gracefully handles anyio cancel scope errors.
|
||||
async def test_mcp_tool_close_cleans_up_in_original_task(caplog):
|
||||
"""Closing an MCP tool from another task should still unwind contexts in the owner task."""
|
||||
import asyncio
|
||||
|
||||
This tests the fix for the bug where calling connect(reset=True) from a
|
||||
different task than where the connection was originally established would
|
||||
cause: RuntimeError: Attempted to exit cancel scope in a different task
|
||||
than it was entered in
|
||||
class TaskBoundTransportContext:
|
||||
def __init__(self) -> None:
|
||||
self.enter_task = None
|
||||
self.exit_task = None
|
||||
self.closed_cleanly = False
|
||||
|
||||
This happens when using multiple MCP tools with AG-UI streaming - the first
|
||||
tool call succeeds, but when the connection closes, the second tool call
|
||||
triggers a reconnection from within the streaming loop (a different task).
|
||||
"""
|
||||
from contextlib import AsyncExitStack
|
||||
async def __aenter__(self):
|
||||
self.enter_task = asyncio.current_task()
|
||||
return (Mock(), Mock())
|
||||
|
||||
from agent_framework._mcp import MCPStdioTool
|
||||
async def __aexit__(self, exc_type, exc, tb):
|
||||
self.exit_task = asyncio.current_task()
|
||||
if self.exit_task is not self.enter_task:
|
||||
raise RuntimeError("Attempted to exit cancel scope in a different task than it was entered in")
|
||||
self.closed_cleanly = True
|
||||
return
|
||||
|
||||
# Use load_tools=False and load_prompts=False to avoid triggering them during connect()
|
||||
tool = MCPStdioTool(
|
||||
tool = MCPStreamableHTTPTool(
|
||||
name="test_server",
|
||||
command="test_command",
|
||||
args=["arg1"],
|
||||
url="https://example.com/mcp",
|
||||
load_tools=False,
|
||||
load_prompts=False,
|
||||
)
|
||||
|
||||
# Mock the exit stack to raise the cross-task cancel scope error
|
||||
mock_exit_stack = AsyncMock(spec=AsyncExitStack)
|
||||
mock_exit_stack.aclose = AsyncMock(
|
||||
side_effect=RuntimeError("Attempted to exit cancel scope in a different task than it was entered in")
|
||||
)
|
||||
tool._exit_stack = mock_exit_stack
|
||||
tool.session = Mock()
|
||||
tool.is_connected = True
|
||||
transport_context = TaskBoundTransportContext()
|
||||
mock_session = Mock()
|
||||
mock_session._request_id = 1
|
||||
mock_session.initialize = AsyncMock()
|
||||
|
||||
# Mock get_mcp_client to return a mock transport
|
||||
mock_transport = (Mock(), Mock())
|
||||
mock_context = AsyncMock()
|
||||
mock_context.__aenter__ = AsyncMock(return_value=mock_transport)
|
||||
mock_context.__aexit__ = AsyncMock()
|
||||
mock_session_context = AsyncMock()
|
||||
mock_session_context.__aenter__ = AsyncMock(return_value=mock_session)
|
||||
mock_session_context.__aexit__ = AsyncMock(return_value=None)
|
||||
|
||||
with (
|
||||
patch.object(tool, "get_mcp_client", return_value=mock_context),
|
||||
patch("agent_framework._mcp.ClientSession") as mock_session_class,
|
||||
patch.object(tool, "get_mcp_client", return_value=transport_context),
|
||||
patch("agent_framework._mcp.ClientSession", return_value=mock_session_context),
|
||||
):
|
||||
mock_session = Mock()
|
||||
mock_session._request_id = 1
|
||||
mock_session.initialize = AsyncMock()
|
||||
mock_session.set_logging_level = AsyncMock()
|
||||
mock_session_context = AsyncMock()
|
||||
mock_session_context.__aenter__ = AsyncMock(return_value=mock_session)
|
||||
mock_session_context.__aexit__ = AsyncMock()
|
||||
mock_session_class.return_value = mock_session_context
|
||||
await asyncio.create_task(tool.connect())
|
||||
|
||||
# This should NOT raise even though aclose() raised the cancel scope error
|
||||
# The _safe_close_exit_stack method should catch and log the error
|
||||
await tool.connect(reset=True)
|
||||
caplog.clear()
|
||||
with caplog.at_level(logging.WARNING, logger=logger.name):
|
||||
await tool.close()
|
||||
|
||||
# Verify a new exit stack was created (the old mock was replaced)
|
||||
assert tool._exit_stack is not mock_exit_stack
|
||||
assert tool.session is not None
|
||||
assert transport_context.closed_cleanly is True
|
||||
assert transport_context.exit_task is transport_context.enter_task
|
||||
assert not any("cancel scope" in record.getMessage().lower() for record in caplog.records)
|
||||
|
||||
|
||||
async def test_mcp_tool_connect_reset_cleans_up_in_original_task(caplog):
|
||||
"""Resetting an MCP tool from another task should unwind and reconnect on the owner task."""
|
||||
import asyncio
|
||||
|
||||
class TaskBoundTransportContext:
|
||||
def __init__(self) -> None:
|
||||
self.enter_task = None
|
||||
self.exit_task = None
|
||||
self.closed_cleanly = False
|
||||
|
||||
async def __aenter__(self):
|
||||
self.enter_task = asyncio.current_task()
|
||||
return (Mock(), Mock())
|
||||
|
||||
async def __aexit__(self, exc_type, exc, tb):
|
||||
self.exit_task = asyncio.current_task()
|
||||
if self.exit_task is not self.enter_task:
|
||||
raise RuntimeError("Attempted to exit cancel scope in a different task than it was entered in")
|
||||
self.closed_cleanly = True
|
||||
return
|
||||
|
||||
tool = MCPStreamableHTTPTool(
|
||||
name="test_server",
|
||||
url="https://example.com/mcp",
|
||||
load_tools=False,
|
||||
load_prompts=False,
|
||||
)
|
||||
|
||||
transport_contexts = [TaskBoundTransportContext(), TaskBoundTransportContext()]
|
||||
sessions = []
|
||||
session_contexts = []
|
||||
for _ in range(2):
|
||||
session = Mock()
|
||||
session._request_id = 1
|
||||
session.initialize = AsyncMock()
|
||||
session.set_logging_level = AsyncMock()
|
||||
sessions.append(session)
|
||||
|
||||
session_context = AsyncMock()
|
||||
session_context.__aenter__ = AsyncMock(return_value=session)
|
||||
session_context.__aexit__ = AsyncMock(return_value=None)
|
||||
session_contexts.append(session_context)
|
||||
|
||||
with (
|
||||
patch.object(tool, "get_mcp_client", side_effect=transport_contexts),
|
||||
patch("agent_framework._mcp.ClientSession", side_effect=session_contexts),
|
||||
):
|
||||
await tool.connect()
|
||||
|
||||
caplog.clear()
|
||||
with caplog.at_level(logging.WARNING, logger=logger.name):
|
||||
await asyncio.create_task(tool.connect(reset=True))
|
||||
|
||||
assert transport_contexts[0].closed_cleanly is True
|
||||
assert transport_contexts[0].exit_task is transport_contexts[0].enter_task
|
||||
assert transport_contexts[1].enter_task is transport_contexts[0].enter_task
|
||||
assert tool.session is sessions[1]
|
||||
assert tool.is_connected is True
|
||||
assert not any("cancel scope" in record.getMessage().lower() for record in caplog.records)
|
||||
|
||||
await tool.close()
|
||||
|
||||
|
||||
async def test_mcp_tool_connect_from_lifecycle_owner_bypasses_request_lock() -> None:
|
||||
"""connect(reset=True) should bypass the request queue when already on the owner task."""
|
||||
import asyncio
|
||||
|
||||
tool = MCPStreamableHTTPTool(
|
||||
name="test_server",
|
||||
url="https://example.com/mcp",
|
||||
load_tools=False,
|
||||
load_prompts=False,
|
||||
)
|
||||
|
||||
async def connect_from_owner_task() -> None:
|
||||
tool._lifecycle_owner_task = asyncio.current_task()
|
||||
try:
|
||||
async with tool._lifecycle_request_lock:
|
||||
await tool.connect(reset=True)
|
||||
finally:
|
||||
tool._lifecycle_owner_task = None
|
||||
|
||||
with patch.object(tool, "_connect_on_owner", AsyncMock()) as mock_connect_on_owner:
|
||||
await asyncio.wait_for(connect_from_owner_task(), timeout=0.1)
|
||||
|
||||
mock_connect_on_owner.assert_awaited_once_with(reset=True)
|
||||
|
||||
|
||||
async def test_mcp_tool_close_from_lifecycle_owner_bypasses_request_lock() -> None:
|
||||
"""close() should bypass the request queue when already on the owner task."""
|
||||
import asyncio
|
||||
|
||||
tool = MCPStreamableHTTPTool(
|
||||
name="test_server",
|
||||
url="https://example.com/mcp",
|
||||
load_tools=False,
|
||||
load_prompts=False,
|
||||
)
|
||||
|
||||
async def close_from_owner_task() -> None:
|
||||
tool._lifecycle_owner_task = asyncio.current_task()
|
||||
try:
|
||||
async with tool._lifecycle_request_lock:
|
||||
await tool.close()
|
||||
finally:
|
||||
tool._lifecycle_owner_task = None
|
||||
|
||||
with patch.object(tool, "_close_on_owner", AsyncMock()) as mock_close_on_owner:
|
||||
await asyncio.wait_for(close_from_owner_task(), timeout=0.1)
|
||||
|
||||
mock_close_on_owner.assert_awaited_once_with()
|
||||
|
||||
|
||||
async def test_mcp_tool_safe_close_reraises_other_runtime_errors():
|
||||
|
||||
@@ -192,10 +192,10 @@ class ConcreteHistoryProvider(BaseHistoryProvider):
|
||||
self.stored: list[Message] = []
|
||||
self._stored_messages = stored_messages or []
|
||||
|
||||
async def get_messages(self, session_id: str | None, **kwargs) -> list[Message]:
|
||||
async def get_messages(self, session_id: str | None, *, state=None, **kwargs) -> list[Message]:
|
||||
return list(self._stored_messages)
|
||||
|
||||
async def save_messages(self, session_id: str | None, messages: Sequence[Message], **kwargs) -> None:
|
||||
async def save_messages(self, session_id: str | None, messages: Sequence[Message], *, state=None, **kwargs) -> None:
|
||||
self.stored.extend(messages)
|
||||
|
||||
|
||||
|
||||
@@ -12,6 +12,7 @@ from agent_framework import (
|
||||
FunctionTool,
|
||||
tool,
|
||||
)
|
||||
from agent_framework._middleware import FunctionInvocationContext
|
||||
from agent_framework._tools import (
|
||||
_parse_annotation,
|
||||
_parse_inputs,
|
||||
@@ -952,6 +953,128 @@ async def test_ai_function_with_kwargs_injection():
|
||||
assert result_default[0].text == "x=10, user=unknown"
|
||||
|
||||
|
||||
async def test_ai_function_with_explicit_invocation_context():
|
||||
"""Test that invoke() can receive runtime kwargs via FunctionInvocationContext."""
|
||||
|
||||
@tool
|
||||
def tool_with_context(x: int, ctx: FunctionInvocationContext) -> str:
|
||||
"""A tool that accepts runtime context injection."""
|
||||
user_id = ctx.kwargs.get("user_id", "unknown")
|
||||
return f"x={x}, user={user_id}"
|
||||
|
||||
assert tool_with_context.parameters() == {
|
||||
"properties": {"x": {"title": "X", "type": "integer"}},
|
||||
"required": ["x"],
|
||||
"title": "tool_with_context_input",
|
||||
"type": "object",
|
||||
}
|
||||
|
||||
context = FunctionInvocationContext(
|
||||
function=tool_with_context,
|
||||
arguments=tool_with_context.input_model(x=7),
|
||||
kwargs={"user_id": "ctx-user"},
|
||||
)
|
||||
|
||||
result = await tool_with_context.invoke(context=context)
|
||||
|
||||
assert result[0].text == "x=7, user=ctx-user"
|
||||
|
||||
|
||||
async def test_ai_function_with_typed_context_parameter_using_custom_name():
|
||||
"""Test that typed context injection works for names other than ctx."""
|
||||
|
||||
@tool
|
||||
def tool_with_runtime_context(x: int, runtime: FunctionInvocationContext) -> str:
|
||||
"""A tool that uses a custom context parameter name."""
|
||||
user_id = runtime.kwargs.get("user_id", "unknown")
|
||||
return f"x={x}, user={user_id}"
|
||||
|
||||
assert tool_with_runtime_context.parameters() == {
|
||||
"properties": {"x": {"title": "X", "type": "integer"}},
|
||||
"required": ["x"],
|
||||
"title": "tool_with_runtime_context_input",
|
||||
"type": "object",
|
||||
}
|
||||
|
||||
context = FunctionInvocationContext(
|
||||
function=tool_with_runtime_context,
|
||||
arguments=tool_with_runtime_context.input_model(x=8),
|
||||
kwargs={"user_id": "runtime-user"},
|
||||
)
|
||||
|
||||
result = await tool_with_runtime_context.invoke(context=context)
|
||||
|
||||
assert result[0].text == "x=8, user=runtime-user"
|
||||
|
||||
|
||||
async def test_ai_function_with_explicit_schema_and_untyped_ctx():
|
||||
"""Test that explicit schemas allow an untyped ctx parameter."""
|
||||
|
||||
class ToolInput(BaseModel):
|
||||
x: int
|
||||
|
||||
@tool(schema=ToolInput)
|
||||
def tool_with_schema(x, ctx) -> str:
|
||||
"""A tool with explicit schema and implicit ctx injection."""
|
||||
return f"x={x}, user={ctx.kwargs.get('user_id', 'unknown')}"
|
||||
|
||||
context = FunctionInvocationContext(
|
||||
function=tool_with_schema,
|
||||
arguments=ToolInput(x=9),
|
||||
kwargs={"user_id": "schema-user"},
|
||||
)
|
||||
|
||||
result = await tool_with_schema.invoke(context=context)
|
||||
|
||||
assert result[0].text == "x=9, user=schema-user"
|
||||
|
||||
|
||||
async def test_ai_function_with_explicit_schema_and_typed_ctx():
|
||||
"""Test that explicit schemas also work with typed context injection."""
|
||||
|
||||
class ToolInput(BaseModel):
|
||||
x: int
|
||||
|
||||
@tool(schema=ToolInput)
|
||||
def tool_with_schema(x: int, runtime: FunctionInvocationContext) -> str:
|
||||
"""A tool with explicit schema and typed context injection."""
|
||||
return f"x={x}, user={runtime.kwargs.get('user_id', 'unknown')}"
|
||||
|
||||
context = FunctionInvocationContext(
|
||||
function=tool_with_schema,
|
||||
arguments=ToolInput(x=11),
|
||||
kwargs={"user_id": "typed-schema-user"},
|
||||
)
|
||||
|
||||
result = await tool_with_schema.invoke(context=context)
|
||||
|
||||
assert tool_with_schema.parameters() == ToolInput.model_json_schema()
|
||||
assert result[0].text == "x=11, user=typed-schema-user"
|
||||
|
||||
|
||||
def test_ai_function_with_multiple_typed_context_parameters_fails():
|
||||
"""Test that tools reject multiple typed FunctionInvocationContext parameters."""
|
||||
|
||||
with pytest.raises(ValueError, match="multiple FunctionInvocationContext parameters"):
|
||||
|
||||
@tool
|
||||
def invalid_tool(ctx_one: FunctionInvocationContext, ctx_two: FunctionInvocationContext) -> str:
|
||||
return f"{ctx_one.kwargs}-{ctx_two.kwargs}"
|
||||
|
||||
|
||||
def test_ai_function_with_ctx_and_typed_context_parameter_fails():
|
||||
"""Test that explicit-schema tools reject both implicit ctx and typed context parameters."""
|
||||
|
||||
class ToolInput(BaseModel):
|
||||
x: int
|
||||
|
||||
with pytest.raises(ValueError, match="multiple FunctionInvocationContext parameters"):
|
||||
|
||||
@tool(schema=ToolInput)
|
||||
def invalid_tool(x, ctx, runtime: FunctionInvocationContext) -> str:
|
||||
return f"{x}-{ctx.kwargs}-{runtime.kwargs}"
|
||||
|
||||
|
||||
# region _parse_annotation tests
|
||||
|
||||
|
||||
|
||||
@@ -10,7 +10,6 @@ from openai.types import CreateEmbeddingResponse
|
||||
from openai.types import Embedding as OpenAIEmbedding
|
||||
from openai.types.create_embedding_response import Usage
|
||||
|
||||
from agent_framework.azure import AzureOpenAIEmbeddingClient
|
||||
from agent_framework.openai import (
|
||||
OpenAIEmbeddingClient,
|
||||
OpenAIEmbeddingOptions,
|
||||
@@ -190,73 +189,6 @@ async def test_openai_empty_values_returns_empty(openai_unit_test_env: None) ->
|
||||
client.client.embeddings.create.assert_not_called()
|
||||
|
||||
|
||||
# --- Azure OpenAI unit tests ---
|
||||
|
||||
|
||||
def test_azure_construction_with_deployment_name() -> None:
|
||||
client = AzureOpenAIEmbeddingClient(
|
||||
deployment_name="text-embedding-3-small",
|
||||
api_key="test-key",
|
||||
endpoint="https://test.openai.azure.com/",
|
||||
)
|
||||
assert client.model_id == "text-embedding-3-small"
|
||||
|
||||
|
||||
def test_azure_construction_with_existing_client() -> None:
|
||||
mock_client = MagicMock()
|
||||
client = AzureOpenAIEmbeddingClient(
|
||||
deployment_name="my-deployment",
|
||||
async_client=mock_client,
|
||||
)
|
||||
assert client.model_id == "my-deployment"
|
||||
assert client.client is mock_client
|
||||
|
||||
|
||||
def test_azure_construction_missing_deployment_name_raises(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.delenv("AZURE_OPENAI_EMBEDDING_DEPLOYMENT_NAME", raising=False)
|
||||
with pytest.raises(ValueError, match="deployment name is required"):
|
||||
AzureOpenAIEmbeddingClient(
|
||||
api_key="test-key",
|
||||
endpoint="https://test.openai.azure.com/",
|
||||
)
|
||||
|
||||
|
||||
def test_azure_construction_missing_credentials_raises() -> None:
|
||||
with pytest.raises(ValueError, match="api_key, credential, or a client"):
|
||||
AzureOpenAIEmbeddingClient(
|
||||
deployment_name="test",
|
||||
endpoint="https://test.openai.azure.com/",
|
||||
)
|
||||
|
||||
|
||||
async def test_azure_get_embeddings() -> None:
|
||||
mock_response = _make_openai_response(
|
||||
embeddings=[[0.1, 0.2]],
|
||||
)
|
||||
mock_async_client = MagicMock()
|
||||
mock_async_client.embeddings = MagicMock()
|
||||
mock_async_client.embeddings.create = AsyncMock(return_value=mock_response)
|
||||
|
||||
client = AzureOpenAIEmbeddingClient(
|
||||
deployment_name="text-embedding-3-small",
|
||||
async_client=mock_async_client,
|
||||
)
|
||||
|
||||
result = await client.get_embeddings(["hello"])
|
||||
|
||||
assert len(result) == 1
|
||||
assert result[0].vector == [0.1, 0.2]
|
||||
|
||||
|
||||
def test_azure_otel_provider_name() -> None:
|
||||
mock_client = MagicMock()
|
||||
client = AzureOpenAIEmbeddingClient(
|
||||
deployment_name="test",
|
||||
async_client=mock_client,
|
||||
)
|
||||
assert client.OTEL_PROVIDER_NAME == "azure.ai.openai"
|
||||
|
||||
|
||||
# --- Integration tests ---
|
||||
|
||||
skip_if_openai_integration_tests_disabled = pytest.mark.skipif(
|
||||
@@ -264,12 +196,6 @@ skip_if_openai_integration_tests_disabled = pytest.mark.skipif(
|
||||
reason="No real OPENAI_API_KEY provided; skipping integration tests.",
|
||||
)
|
||||
|
||||
skip_if_azure_openai_integration_tests_disabled = pytest.mark.skipif(
|
||||
not os.getenv("AZURE_OPENAI_ENDPOINT")
|
||||
or (not os.getenv("AZURE_OPENAI_API_KEY") and not os.getenv("AZURE_OPENAI_EMBEDDING_DEPLOYMENT_NAME")),
|
||||
reason="No Azure OpenAI credentials provided; skipping integration tests.",
|
||||
)
|
||||
|
||||
|
||||
@skip_if_openai_integration_tests_disabled
|
||||
@pytest.mark.flaky
|
||||
@@ -315,49 +241,3 @@ async def test_integration_openai_get_embeddings_with_dimensions() -> None:
|
||||
|
||||
assert len(result) == 1
|
||||
assert len(result[0].vector) == 256
|
||||
|
||||
|
||||
@skip_if_azure_openai_integration_tests_disabled
|
||||
@pytest.mark.flaky
|
||||
@pytest.mark.integration
|
||||
async def test_integration_azure_openai_get_embeddings() -> None:
|
||||
"""End-to-end test of Azure OpenAI embedding generation."""
|
||||
client = AzureOpenAIEmbeddingClient()
|
||||
|
||||
result = await client.get_embeddings(["hello world"])
|
||||
|
||||
assert len(result) == 1
|
||||
assert isinstance(result[0].vector, list)
|
||||
assert len(result[0].vector) > 0
|
||||
assert all(isinstance(v, float) for v in result[0].vector)
|
||||
assert result[0].model_id is not None
|
||||
assert result.usage is not None
|
||||
assert result.usage["input_token_count"] > 0
|
||||
|
||||
|
||||
@skip_if_azure_openai_integration_tests_disabled
|
||||
@pytest.mark.flaky
|
||||
@pytest.mark.integration
|
||||
async def test_integration_azure_openai_get_embeddings_multiple() -> None:
|
||||
"""Test Azure OpenAI embedding generation for multiple inputs."""
|
||||
client = AzureOpenAIEmbeddingClient()
|
||||
|
||||
result = await client.get_embeddings(["hello", "world", "test"])
|
||||
|
||||
assert len(result) == 3
|
||||
dims = [len(e.vector) for e in result]
|
||||
assert all(d == dims[0] for d in dims)
|
||||
|
||||
|
||||
@skip_if_azure_openai_integration_tests_disabled
|
||||
@pytest.mark.flaky
|
||||
@pytest.mark.integration
|
||||
async def test_integration_azure_openai_get_embeddings_with_dimensions() -> None:
|
||||
"""Test Azure OpenAI embedding generation with custom dimensions."""
|
||||
client = AzureOpenAIEmbeddingClient()
|
||||
|
||||
options: OpenAIEmbeddingOptions = {"dimensions": 256}
|
||||
result = await client.get_embeddings(["hello world"], options=options)
|
||||
|
||||
assert len(result) == 1
|
||||
assert len(result[0].vector) == 256
|
||||
|
||||
@@ -1876,6 +1876,19 @@ def test_prepare_tools_for_openai_with_image_generation_options() -> None:
|
||||
assert image_tool["quality"] == "high"
|
||||
|
||||
|
||||
def test_prepare_tools_for_openai_with_custom_image_generation_model() -> None:
|
||||
"""Test image generation tool conversion with a custom model string."""
|
||||
client = OpenAIResponsesClient(model_id="test-model", api_key="test-key")
|
||||
|
||||
tool = OpenAIResponsesClient.get_image_generation_tool(model="custom-image-model")
|
||||
|
||||
resp_tools = client._prepare_tools_for_openai([tool])
|
||||
assert len(resp_tools) == 1
|
||||
image_tool = resp_tools[0]
|
||||
assert image_tool["type"] == "image_generation"
|
||||
assert image_tool["model"] == "custom-image-model"
|
||||
|
||||
|
||||
def test_parse_chunk_from_openai_with_mcp_approval_request() -> None:
|
||||
"""Test that a streaming mcp_approval_request event is parsed into FunctionApprovalRequestContent."""
|
||||
client = OpenAIResponsesClient(model_id="test-model", api_key="test-key")
|
||||
|
||||
@@ -13,6 +13,8 @@ from agent_framework import (
|
||||
AgentRunInputs,
|
||||
AgentSession,
|
||||
BaseAgent,
|
||||
Case,
|
||||
Default,
|
||||
Executor,
|
||||
Message,
|
||||
ResponseStream,
|
||||
@@ -223,6 +225,29 @@ def test_add_edge_with_condition():
|
||||
assert "Target" in workflow.executors
|
||||
|
||||
|
||||
def test_switch_case_with_agents():
|
||||
"""Test add_switch_case_edge_group with Case and Default edges using agents."""
|
||||
router = DummyAgent(id="router_agent", name="router")
|
||||
handler = DummyAgent(id="handler", name="handler")
|
||||
fallback = DummyAgent(id="fallback_agent", name="fallback")
|
||||
|
||||
workflow = (
|
||||
WorkflowBuilder(start_executor=router)
|
||||
.add_switch_case_edge_group(
|
||||
router,
|
||||
[
|
||||
Case(condition=lambda _: True, target=handler),
|
||||
Default(target=fallback),
|
||||
],
|
||||
)
|
||||
.build()
|
||||
)
|
||||
|
||||
# All three agents should be AgentExecutor wrappers
|
||||
agent_executors = [e for e in workflow.executors.values() if isinstance(e, AgentExecutor)]
|
||||
assert len(agent_executors) == 3
|
||||
|
||||
|
||||
# region with_output_from tests
|
||||
|
||||
|
||||
|
||||
@@ -23,12 +23,12 @@ classifiers = [
|
||||
]
|
||||
dependencies = [
|
||||
"agent-framework-core>=1.0.0rc4",
|
||||
"powerfx>=0.0.31; python_version < '3.14'",
|
||||
"powerfx>=0.0.32,<0.0.35; python_version < '3.14'",
|
||||
"pyyaml>=6.0,<7.0",
|
||||
]
|
||||
[dependency-groups]
|
||||
dev = [
|
||||
"types-PyYaml"
|
||||
"types-PyYaml==6.0.12.20250915"
|
||||
]
|
||||
|
||||
[tool.uv]
|
||||
@@ -94,7 +94,7 @@ include = "../../shared_tasks.toml"
|
||||
|
||||
[tool.poe.tasks]
|
||||
mypy = "mypy --config-file $POE_ROOT/pyproject.toml agent_framework_declarative"
|
||||
test = "pytest -m \"not integration\" --cov=agent_framework_declarative --cov-report=term-missing:skip-covered tests"
|
||||
test = 'pytest -m "not integration" --cov=agent_framework_declarative --cov-report=term-missing:skip-covered tests'
|
||||
|
||||
[build-system]
|
||||
requires = ["flit-core >= 3.11,<4.0"]
|
||||
|
||||
+5
-3
@@ -2509,6 +2509,8 @@
|
||||
},
|
||||
"node_modules/@tailwindcss/oxide-wasm32-wasi/node_modules/@napi-rs/wasm-runtime": {
|
||||
"version": "0.2.12",
|
||||
"resolved": "https://registry.npmjs.org/@napi-rs/wasm-runtime/-/wasm-runtime-0.2.12.tgz",
|
||||
"integrity": "sha512-ZVWUcfwY4E/yPitQJl481FjFo3K22D6qF0DuFH6Y/nbnE11GY5uguDxZMGXPQ8WQ0128MXQD7TnfHyK4oWoIJQ==",
|
||||
"inBundle": true,
|
||||
"license": "MIT",
|
||||
"optional": true,
|
||||
@@ -5019,9 +5021,9 @@
|
||||
}
|
||||
},
|
||||
"node_modules/tar": {
|
||||
"version": "7.5.9",
|
||||
"resolved": "https://registry.npmjs.org/tar/-/tar-7.5.9.tgz",
|
||||
"integrity": "sha512-BTLcK0xsDh2+PUe9F6c2TlRp4zOOBMTkoQHQIWSIzI0R7KG46uEwq4OPk2W7bZcprBMsuaeFsqwYr7pjh6CuHg==",
|
||||
"version": "7.5.11",
|
||||
"resolved": "https://registry.npmjs.org/tar/-/tar-7.5.11.tgz",
|
||||
"integrity": "sha512-ChjMH33/KetonMTAtpYdgUFr0tbz69Fp2v7zWxQfYZX4g5ZN2nOBXm1R2xyA+lMIKrLKIoKAwFj93jE/avX9cQ==",
|
||||
"license": "BlueOak-1.0.0",
|
||||
"dependencies": {
|
||||
"@isaacs/fs-minipass": "^4.0.0",
|
||||
|
||||
@@ -169,24 +169,24 @@
|
||||
"@babel/helper-validator-identifier" "^7.27.1"
|
||||
|
||||
"@emnapi/core@^1.4.3", "@emnapi/core@^1.4.5":
|
||||
version "1.8.1"
|
||||
resolved "https://registry.yarnpkg.com/@emnapi/core/-/core-1.8.1.tgz#fd9efe721a616288345ffee17a1f26ac5dd01349"
|
||||
integrity sha512-AvT9QFpxK0Zd8J0jopedNm+w/2fIzvtPKPjqyw9jwvBaReTTqPBk9Hixaz7KbjimP+QNz605/XnjFcDAL2pqBg==
|
||||
version "1.9.0"
|
||||
resolved "https://registry.yarnpkg.com/@emnapi/core/-/core-1.9.0.tgz#4a54213b208fcf288cce25076c74e0f7613e6100"
|
||||
integrity sha512-0DQ98G9ZQZOxfUcQn1waV2yS8aWdZ6kJMbYCJB3oUBecjWYO1fqJ+a1DRfPF3O5JEkwqwP1A9QEN/9mYm2Yd0w==
|
||||
dependencies:
|
||||
"@emnapi/wasi-threads" "1.1.0"
|
||||
"@emnapi/wasi-threads" "1.2.0"
|
||||
tslib "^2.4.0"
|
||||
|
||||
"@emnapi/runtime@^1.4.3", "@emnapi/runtime@^1.4.5":
|
||||
version "1.8.1"
|
||||
resolved "https://registry.yarnpkg.com/@emnapi/runtime/-/runtime-1.8.1.tgz#550fa7e3c0d49c5fb175a116e8cd70614f9a22a5"
|
||||
integrity sha512-mehfKSMWjjNol8659Z8KxEMrdSJDDot5SXMq00dM8BN4o+CLNXQ0xH2V7EchNHV4RmbZLmmPdEaXZc5H2FXmDg==
|
||||
version "1.9.0"
|
||||
resolved "https://registry.yarnpkg.com/@emnapi/runtime/-/runtime-1.9.0.tgz#91c54a6e77c36154c125e873409472e2b70efd5b"
|
||||
integrity sha512-QN75eB0IH2ywSpRpNddCRfQIhmJYBCJ1x5Lb3IscKAL8bMnVAKnRg8dCoXbHzVLLH7P38N2Z3mtulB7W0J0FKw==
|
||||
dependencies:
|
||||
tslib "^2.4.0"
|
||||
|
||||
"@emnapi/wasi-threads@1.1.0", "@emnapi/wasi-threads@^1.0.4":
|
||||
version "1.1.0"
|
||||
resolved "https://registry.yarnpkg.com/@emnapi/wasi-threads/-/wasi-threads-1.1.0.tgz#60b2102fddc9ccb78607e4a3cf8403ea69be41bf"
|
||||
integrity sha512-WI0DdZ8xFSbgMjR1sFsKABJ/C5OnRrjT06JXbZKexJGrDuPTzZdDYfFlsgcCXCyf+suG5QU2e/y1Wo2V/OapLQ==
|
||||
"@emnapi/wasi-threads@1.2.0", "@emnapi/wasi-threads@^1.0.4":
|
||||
version "1.2.0"
|
||||
resolved "https://registry.yarnpkg.com/@emnapi/wasi-threads/-/wasi-threads-1.2.0.tgz#a19d9772cc3d195370bf6e2a805eec40aa75e18e"
|
||||
integrity sha512-N10dEJNSsUx41Z6pZsXU8FjPjpBEplgH24sfkmITrBED1/U2Esum9F3lfLrMjKHHjmi557zQn7kR9R+XWXu5Rg==
|
||||
dependencies:
|
||||
tslib "^2.4.0"
|
||||
|
||||
@@ -212,7 +212,7 @@
|
||||
|
||||
"@esbuild/darwin-arm64@0.25.9":
|
||||
version "0.25.9"
|
||||
resolved "https://registry.npmjs.org/@esbuild/darwin-arm64/-/darwin-arm64-0.25.9.tgz"
|
||||
resolved "https://registry.yarnpkg.com/@esbuild/darwin-arm64/-/darwin-arm64-0.25.9.tgz#f1513eaf9ec8fa15dcaf4c341b0f005d3e8b47ae"
|
||||
integrity sha512-XIpIDMAjOELi/9PB30vEbVMs3GV1v2zkkPnuyRRURbhqjyzIINwj+nbQATh4H9GxUgH1kFsEyQMxwiLFKUS6Rg==
|
||||
|
||||
"@esbuild/darwin-x64@0.25.9":
|
||||
@@ -317,7 +317,7 @@
|
||||
|
||||
"@esbuild/win32-x64@0.25.9":
|
||||
version "0.25.9"
|
||||
resolved "https://registry.yarnpkg.com/@esbuild/win32-x64/-/win32-x64-0.25.9.tgz#585624dc829cfb6e7c0aa6c3ca7d7e6daa87e34f"
|
||||
resolved "https://registry.npmjs.org/@esbuild/win32-x64/-/win32-x64-0.25.9.tgz"
|
||||
integrity sha512-PPOl1mi6lpLNQxnGoyAfschAodRFYXJ+9fs6WHXz7CSWKbOqiMZsubC+BQsVKuul+3vKLuwTHsS2c2y9EoKwxQ==
|
||||
|
||||
"@eslint-community/eslint-utils@^4.2.0", "@eslint-community/eslint-utils@^4.7.0":
|
||||
@@ -984,12 +984,12 @@
|
||||
|
||||
"@rollup/rollup-win32-x64-gnu@4.59.0":
|
||||
version "4.59.0"
|
||||
resolved "https://registry.yarnpkg.com/@rollup/rollup-win32-x64-gnu/-/rollup-win32-x64-gnu-4.59.0.tgz#c4af3e9518c9a5cd4b1c163dc81d0ad4d82e7eab"
|
||||
resolved "https://registry.npmjs.org/@rollup/rollup-win32-x64-gnu/-/rollup-win32-x64-gnu-4.59.0.tgz"
|
||||
integrity sha512-laBkYlSS1n2L8fSo1thDNGrCTQMmxjYY5G0WFWjFFYZkKPjsMBsgJfGf4TLxXrF6RyhI60L8TMOjBMvXiTcxeA==
|
||||
|
||||
"@rollup/rollup-win32-x64-msvc@4.59.0":
|
||||
version "4.59.0"
|
||||
resolved "https://registry.yarnpkg.com/@rollup/rollup-win32-x64-msvc/-/rollup-win32-x64-msvc-4.59.0.tgz#4584a8a87b29188a4c1fe987a9fcf701e256d86c"
|
||||
resolved "https://registry.npmjs.org/@rollup/rollup-win32-x64-msvc/-/rollup-win32-x64-msvc-4.59.0.tgz"
|
||||
integrity sha512-2HRCml6OztYXyJXAvdDXPKcawukWY2GpR5/nxKp4iBgiO3wcoEGkAaqctIbZcNB6KlUQBIqt8VYkNSj2397EfA==
|
||||
|
||||
"@tailwindcss/node@4.1.12":
|
||||
@@ -1012,7 +1012,7 @@
|
||||
|
||||
"@tailwindcss/oxide-darwin-arm64@4.1.12":
|
||||
version "4.1.12"
|
||||
resolved "https://registry.npmjs.org/@tailwindcss/oxide-darwin-arm64/-/oxide-darwin-arm64-4.1.12.tgz"
|
||||
resolved "https://registry.yarnpkg.com/@tailwindcss/oxide-darwin-arm64/-/oxide-darwin-arm64-4.1.12.tgz#e8bd4798f26ec1d012bf0683aeb77449f71505cd"
|
||||
integrity sha512-cq1qmq2HEtDV9HvZlTtrj671mCdGB93bVY6J29mwCyaMYCP/JaUBXxrQQQm7Qn33AXXASPUb2HFZlWiiHWFytw==
|
||||
|
||||
"@tailwindcss/oxide-darwin-x64@4.1.12":
|
||||
@@ -1069,7 +1069,7 @@
|
||||
|
||||
"@tailwindcss/oxide-win32-x64-msvc@4.1.12":
|
||||
version "4.1.12"
|
||||
resolved "https://registry.yarnpkg.com/@tailwindcss/oxide-win32-x64-msvc/-/oxide-win32-x64-msvc-4.1.12.tgz#b1ee2ed0ef2c4095ddec3684a1987e2b3613af36"
|
||||
resolved "https://registry.npmjs.org/@tailwindcss/oxide-win32-x64-msvc/-/oxide-win32-x64-msvc-4.1.12.tgz"
|
||||
integrity sha512-NKIh5rzw6CpEodv/++r0hGLlfgT/gFN+5WNdZtvh6wpU2BpGNgdjvj6H2oFc8nCM839QM1YOhjpgbAONUb4IxA==
|
||||
|
||||
"@tailwindcss/oxide@4.1.12":
|
||||
@@ -1398,7 +1398,7 @@ brace-expansion@^1.1.7:
|
||||
|
||||
brace-expansion@^2.0.2:
|
||||
version "2.0.2"
|
||||
resolved "https://registry.yarnpkg.com/brace-expansion/-/brace-expansion-2.0.2.tgz#54fc53237a613d854c7bd37463aad17df87214e7"
|
||||
resolved "https://registry.npmjs.org/brace-expansion/-/brace-expansion-2.0.2.tgz"
|
||||
integrity sha512-Jt0vHyM+jmUBqojB7E1NIYadt0vI0Qxjxd2TErW94wDz+E2LAm5vKMXXwg6ZZBTHPuUlDgQHKXvjGBdfcF1ZDQ==
|
||||
dependencies:
|
||||
balanced-match "^1.0.0"
|
||||
@@ -1812,7 +1812,7 @@ flatted@^3.2.9:
|
||||
|
||||
fsevents@~2.3.2, fsevents@~2.3.3:
|
||||
version "2.3.3"
|
||||
resolved "https://registry.npmjs.org/fsevents/-/fsevents-2.3.3.tgz"
|
||||
resolved "https://registry.yarnpkg.com/fsevents/-/fsevents-2.3.3.tgz#cac6407785d03675a2a5e1a5305c697b347d90d6"
|
||||
integrity sha512-5xoDfX+fL7faATnagmWPpbFtwh/R77WmMMqqHGS65C3vvB0YHrgF+B1YmZ3441tMj5n63k0212XNoJwzlhffQw==
|
||||
|
||||
gensync@^1.0.0-beta.2:
|
||||
@@ -1921,7 +1921,7 @@ js-tokens@^4.0.0:
|
||||
|
||||
js-yaml@^4.1.0:
|
||||
version "4.1.1"
|
||||
resolved "https://registry.yarnpkg.com/js-yaml/-/js-yaml-4.1.1.tgz#854c292467705b699476e1a2decc0c8a3458806b"
|
||||
resolved "https://registry.npmjs.org/js-yaml/-/js-yaml-4.1.1.tgz"
|
||||
integrity sha512-qQKT4zQxXl8lLwBtHMWwaTcGfFOZviOJet3Oy/xmGk2gZH677CJM9EvtfdSkgWcATZhj/55JZ0rmy3myCT5lsA==
|
||||
dependencies:
|
||||
argparse "^2.0.1"
|
||||
@@ -1968,7 +1968,7 @@ levn@^0.4.1:
|
||||
|
||||
lightningcss-darwin-arm64@1.30.1:
|
||||
version "1.30.1"
|
||||
resolved "https://registry.npmjs.org/lightningcss-darwin-arm64/-/lightningcss-darwin-arm64-1.30.1.tgz"
|
||||
resolved "https://registry.yarnpkg.com/lightningcss-darwin-arm64/-/lightningcss-darwin-arm64-1.30.1.tgz#3d47ce5e221b9567c703950edf2529ca4a3700ae"
|
||||
integrity sha512-c8JK7hyE65X1MHMN+Viq9n11RRC7hgin3HhYKhrMyaXflk5GVplZ60IxyoVtzILeKr+xAJwg6zK6sjTBJ0FKYQ==
|
||||
|
||||
lightningcss-darwin-x64@1.30.1:
|
||||
@@ -2013,7 +2013,7 @@ lightningcss-win32-arm64-msvc@1.30.1:
|
||||
|
||||
lightningcss-win32-x64-msvc@1.30.1:
|
||||
version "1.30.1"
|
||||
resolved "https://registry.yarnpkg.com/lightningcss-win32-x64-msvc/-/lightningcss-win32-x64-msvc-1.30.1.tgz#fd7dd008ea98494b85d24b4bea016793f2e0e352"
|
||||
resolved "https://registry.npmjs.org/lightningcss-win32-x64-msvc/-/lightningcss-win32-x64-msvc-1.30.1.tgz"
|
||||
integrity sha512-PVqXh48wh4T53F/1CCu8PIPCxLzWyCnn/9T5W1Jpmdy5h9Cwd+0YQS6/LwhHXSafuc61/xg9Lv5OrCby6a++jg==
|
||||
|
||||
lightningcss@1.30.1:
|
||||
@@ -2080,14 +2080,14 @@ micromatch@^4.0.8:
|
||||
|
||||
minimatch@^3.1.2:
|
||||
version "3.1.5"
|
||||
resolved "https://registry.yarnpkg.com/minimatch/-/minimatch-3.1.5.tgz#580c88f8d5445f2bd6aa8f3cadefa0de79fbd69e"
|
||||
resolved "https://registry.npmjs.org/minimatch/-/minimatch-3.1.5.tgz"
|
||||
integrity sha512-VgjWUsnnT6n+NUk6eZq77zeFdpW2LWDzP6zFGrCbHXiYNul5Dzqk2HHQ5uFH2DNW5Xbp8+jVzaeNt94ssEEl4w==
|
||||
dependencies:
|
||||
brace-expansion "^1.1.7"
|
||||
|
||||
minimatch@^9.0.4:
|
||||
version "9.0.9"
|
||||
resolved "https://registry.yarnpkg.com/minimatch/-/minimatch-9.0.9.tgz#9b0cb9fcb78087f6fd7eababe2511c4d3d60574e"
|
||||
resolved "https://registry.npmjs.org/minimatch/-/minimatch-9.0.9.tgz"
|
||||
integrity sha512-OBwBN9AL4dqmETlpS2zasx+vTeWclWzkblfZk7KTA5j3jeOONz/tRCnZomUyvNg83wL5Zv9Ss6HMJXAgL8R2Yg==
|
||||
dependencies:
|
||||
brace-expansion "^2.0.2"
|
||||
@@ -2099,7 +2099,7 @@ minipass@^7.0.4, minipass@^7.1.2:
|
||||
|
||||
minizlib@^3.1.0:
|
||||
version "3.1.0"
|
||||
resolved "https://registry.yarnpkg.com/minizlib/-/minizlib-3.1.0.tgz#6ad76c3a8f10227c9b51d1c9ac8e30b27f5a251c"
|
||||
resolved "https://registry.npmjs.org/minizlib/-/minizlib-3.1.0.tgz"
|
||||
integrity sha512-KZxYo1BUkWD2TVFLr0MQoM8vUUigWD3LlD83a/75BqC+4qE0Hb1Vo5v1FgcfaNXvfXzr+5EhQ6ing/CaBijTlw==
|
||||
dependencies:
|
||||
minipass "^7.1.2"
|
||||
@@ -2267,7 +2267,7 @@ reusify@^1.0.4:
|
||||
|
||||
rollup@^4.43.0:
|
||||
version "4.59.0"
|
||||
resolved "https://registry.yarnpkg.com/rollup/-/rollup-4.59.0.tgz#cf74edac17c1486f562d728a4d923a694abdf06f"
|
||||
resolved "https://registry.npmjs.org/rollup/-/rollup-4.59.0.tgz"
|
||||
integrity sha512-2oMpl67a3zCH9H79LeMcbDhXW/UmWG/y2zuqnF2jQq5uq9TbM9TVyXvA4+t+ne2IIkBdrLpAaRQAvo7YI/Yyeg==
|
||||
dependencies:
|
||||
"@types/estree" "1.0.8"
|
||||
@@ -2366,9 +2366,9 @@ tapable@^2.2.0:
|
||||
integrity sha512-Re10+NauLTMCudc7T5WLFLAwDhQ0JWdrMK+9B2M8zR5hRExKmsRDCBA7/aV/pNJFltmBFO5BAMlQFi/vq3nKOg==
|
||||
|
||||
tar@^7.4.3:
|
||||
version "7.5.9"
|
||||
resolved "https://registry.yarnpkg.com/tar/-/tar-7.5.9.tgz#817ac12a54bc4362c51340875b8985d7dc9724b8"
|
||||
integrity sha512-BTLcK0xsDh2+PUe9F6c2TlRp4zOOBMTkoQHQIWSIzI0R7KG46uEwq4OPk2W7bZcprBMsuaeFsqwYr7pjh6CuHg==
|
||||
version "7.5.11"
|
||||
resolved "https://registry.npmjs.org/tar/-/tar-7.5.11.tgz"
|
||||
integrity sha512-ChjMH33/KetonMTAtpYdgUFr0tbz69Fp2v7zWxQfYZX4g5ZN2nOBXm1R2xyA+lMIKrLKIoKAwFj93jE/avX9cQ==
|
||||
dependencies:
|
||||
"@isaacs/fs-minipass" "^4.0.0"
|
||||
chownr "^3.0.0"
|
||||
|
||||
@@ -24,14 +24,20 @@ classifiers = [
|
||||
]
|
||||
dependencies = [
|
||||
"agent-framework-core>=1.0.0rc4",
|
||||
"fastapi>=0.104.0",
|
||||
"uvicorn[standard]>=0.24.0",
|
||||
"python-dotenv>=1.0.0",
|
||||
"fastapi>=0.115.0,<0.133.1",
|
||||
"uvicorn[standard]>=0.30.0,<0.42.0"
|
||||
]
|
||||
|
||||
[project.optional-dependencies]
|
||||
dev = ["pytest>=7.0.0", "watchdog>=3.0.0", "agent-framework-orchestrations"]
|
||||
all = ["pytest>=7.0.0", "watchdog>=3.0.0"]
|
||||
dev = [
|
||||
"pytest==9.0.2",
|
||||
"watchdog==6.0.0",
|
||||
"agent-framework-orchestrations==1.0.0b260311",
|
||||
]
|
||||
all = [
|
||||
"pytest==9.0.2",
|
||||
"watchdog==6.0.0",
|
||||
]
|
||||
|
||||
[project.scripts]
|
||||
devui = "agent_framework_devui:main"
|
||||
@@ -94,7 +100,7 @@ include = "../../shared_tasks.toml"
|
||||
|
||||
[tool.poe.tasks]
|
||||
mypy = "mypy --config-file $POE_ROOT/pyproject.toml agent_framework_devui"
|
||||
test = "pytest -m \"not integration\" --cov=agent_framework_devui --cov-report=term-missing:skip-covered tests"
|
||||
test = 'pytest -m "not integration" --cov=agent_framework_devui --cov-report=term-missing:skip-covered tests'
|
||||
|
||||
[build-system]
|
||||
requires = ["flit-core >= 3.11,<4.0"]
|
||||
|
||||
@@ -29,18 +29,16 @@ Durable execution support for long-running agent workflows using Azure Durable F
|
||||
## Usage
|
||||
|
||||
```python
|
||||
from durabletask.client import TaskHubGrpcClient
|
||||
from durabletask.worker import TaskHubGrpcWorker
|
||||
from agent_framework import ChatAgent
|
||||
from agent_framework import Agent
|
||||
from agent_framework.azure import AzureOpenAIChatClient
|
||||
from agent_framework_durabletask import DurableAIAgentClient, DurableAIAgentWorker
|
||||
from durabletask.client import TaskHubGrpcClient
|
||||
from durabletask.worker import TaskHubGrpcWorker
|
||||
|
||||
# Client side
|
||||
dt_client = TaskHubGrpcClient(host_address="localhost:4001")
|
||||
agent_client = DurableAIAgentClient(dt_client)
|
||||
agent = agent_client.get_agent("assistant")
|
||||
response = agent.run("Hello, how are you?")
|
||||
print(response.text)
|
||||
durable_agent = agent_client.get_agent("assistant")
|
||||
|
||||
# Worker side
|
||||
dt_worker = TaskHubGrpcWorker(host_address="localhost:4001")
|
||||
@@ -48,10 +46,8 @@ agent_worker = DurableAIAgentWorker(dt_worker)
|
||||
|
||||
# Create a chat client for the agent
|
||||
chat_client = AzureOpenAIChatClient()
|
||||
my_agent = ChatAgent(chat_client=chat_client, name="assistant")
|
||||
my_agent = Agent(client=chat_client, name="assistant")
|
||||
agent_worker.add_agent(my_agent)
|
||||
|
||||
dt_worker.start()
|
||||
```
|
||||
|
||||
## Import Path
|
||||
|
||||
@@ -15,17 +15,18 @@ The durable task integration lets you host Microsoft Agent Framework agents usin
|
||||
### Basic Usage Example
|
||||
|
||||
```python
|
||||
from agent_framework import Agent
|
||||
from agent_framework.azure import AzureOpenAIChatClient
|
||||
from agent_framework_durabletask import DurableAIAgentWorker
|
||||
from durabletask.worker import TaskHubGrpcWorker
|
||||
from agent_framework.azure import DurableAIAgentWorker
|
||||
|
||||
# Create the worker
|
||||
with TaskHubGrpcWorker(...) as worker:
|
||||
worker = TaskHubGrpcWorker(host_address="localhost:4001")
|
||||
agent_worker = DurableAIAgentWorker(worker)
|
||||
|
||||
# Register the agent worker wrapper
|
||||
agent_worker = DurableAIAgentWorker(worker)
|
||||
|
||||
# Register the agent
|
||||
agent_worker.add_agent(my_agent)
|
||||
chat_client = AzureOpenAIChatClient()
|
||||
my_agent = Agent(client=chat_client, name="assistant")
|
||||
agent_worker.add_agent(my_agent)
|
||||
```
|
||||
|
||||
For more details, review the Python [README](https://github.com/microsoft/agent-framework/tree/main/python/README.md) and the samples directory.
|
||||
|
||||
@@ -124,10 +124,20 @@ class DurableAgentExecutor(ABC, Generic[TaskT]):
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
def get_new_session(self, agent_name: str, **kwargs: Any) -> DurableAgentSession:
|
||||
def get_new_session(
|
||||
self,
|
||||
agent_name: str,
|
||||
*,
|
||||
session_id: str | None = None,
|
||||
service_session_id: str | None = None,
|
||||
) -> DurableAgentSession:
|
||||
"""Create a new DurableAgentSession with random session ID."""
|
||||
session_id = self._create_session_id(agent_name)
|
||||
return DurableAgentSession.from_session_id(session_id, **kwargs)
|
||||
durable_session_id = self._create_session_id(agent_name)
|
||||
return DurableAgentSession(
|
||||
durable_session_id=durable_session_id,
|
||||
session_id=session_id,
|
||||
service_session_id=service_session_id,
|
||||
)
|
||||
|
||||
def _create_session_id(
|
||||
self,
|
||||
|
||||
@@ -284,46 +284,48 @@ class DurableAgentSession(AgentSession):
|
||||
durable_session_id: AgentSessionId | None = None,
|
||||
session_id: str | None = None,
|
||||
service_session_id: str | None = None,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
super().__init__(session_id=session_id, service_session_id=service_session_id, **kwargs)
|
||||
self._session_id_value: AgentSessionId | None = durable_session_id
|
||||
super().__init__(session_id=session_id, service_session_id=service_session_id)
|
||||
self.durable_session_id: AgentSessionId | None = durable_session_id
|
||||
|
||||
@property
|
||||
def durable_session_id(self) -> AgentSessionId | None:
|
||||
return self._session_id_value
|
||||
|
||||
@durable_session_id.setter
|
||||
def durable_session_id(self, value: AgentSessionId | None) -> None:
|
||||
self._session_id_value = value
|
||||
def to_dict(self) -> dict[str, Any]:
|
||||
state = super().to_dict()
|
||||
if self.durable_session_id is not None:
|
||||
state[self._SERIALIZED_SESSION_ID_KEY] = str(self.durable_session_id)
|
||||
return state
|
||||
|
||||
@classmethod
|
||||
def from_session_id(
|
||||
cls,
|
||||
session_id: AgentSessionId,
|
||||
**kwargs: Any,
|
||||
durable_session_id: AgentSessionId,
|
||||
*,
|
||||
session_id: str | None = None,
|
||||
service_session_id: str | None = None,
|
||||
) -> DurableAgentSession:
|
||||
return cls(durable_session_id=session_id, **kwargs)
|
||||
|
||||
def to_dict(self) -> dict[str, Any]:
|
||||
state = super().to_dict()
|
||||
if self._session_id_value is not None:
|
||||
state[self._SERIALIZED_SESSION_ID_KEY] = str(self._session_id_value)
|
||||
return state
|
||||
"""Create a DurableAgentSession from an AgentSessionId."""
|
||||
return cls(
|
||||
durable_session_id=durable_session_id,
|
||||
session_id=session_id,
|
||||
service_session_id=service_session_id,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, data: dict[str, Any]) -> DurableAgentSession:
|
||||
state_payload = dict(data)
|
||||
session_id_value = state_payload.pop(cls._SERIALIZED_SESSION_ID_KEY, None)
|
||||
session = super().from_dict(state_payload)
|
||||
"""Create a DurableAgentSession from a state dict."""
|
||||
data = dict(data) # defensive copy — avoid mutating caller's dict
|
||||
session_id_value = data.pop(cls._SERIALIZED_SESSION_ID_KEY, None)
|
||||
session = super().from_dict(data)
|
||||
durable_session_id: AgentSessionId | None = None
|
||||
# We need to create a DurableAgentSession from the base AgentSession
|
||||
if session_id_value is not None:
|
||||
if not isinstance(session_id_value, str):
|
||||
raise ValueError("durable_session_id must be a string when present in serialized state")
|
||||
durable_session_id = AgentSessionId.parse(session_id_value)
|
||||
|
||||
durable_session = cls(
|
||||
durable_session_id=durable_session_id,
|
||||
session_id=session.session_id,
|
||||
service_session_id=session.service_session_id,
|
||||
)
|
||||
durable_session.state.update(session.state)
|
||||
if session_id_value is not None:
|
||||
if not isinstance(session_id_value, str):
|
||||
raise ValueError("durable_session_id must be a string when present in serialized state")
|
||||
durable_session._session_id_value = AgentSessionId.parse(session_id_value)
|
||||
return durable_session
|
||||
|
||||
@@ -133,16 +133,13 @@ class DurableAIAgent(SupportsAgentRun, Generic[TaskT]):
|
||||
session=session,
|
||||
)
|
||||
|
||||
def create_session(self, **kwargs: Any) -> DurableAgentSession:
|
||||
def create_session(self, *, session_id: str | None = None) -> DurableAgentSession:
|
||||
"""Create a new agent session via the provider."""
|
||||
return self._executor.get_new_session(self.name, **kwargs)
|
||||
return self._executor.get_new_session(self.name)
|
||||
|
||||
def get_session(self, **kwargs: Any) -> AgentSession:
|
||||
"""Retrieve an existing session via the provider.
|
||||
|
||||
For durable agents, sessions do not use `service_session_id` so this is not used.
|
||||
"""
|
||||
return self._executor.get_new_session(self.name, **kwargs)
|
||||
def get_session(self, service_session_id: str, *, session_id: str | None = None) -> AgentSession:
|
||||
"""Retrieve an existing session via the provider."""
|
||||
return self._executor.get_new_session(self.name, service_session_id=service_session_id, session_id=session_id)
|
||||
|
||||
def _normalize_messages(self, messages: AgentRunInputs | None) -> str:
|
||||
"""Convert supported message inputs to a single string.
|
||||
|
||||
@@ -29,9 +29,10 @@ class DurableAIAgentWorker:
|
||||
|
||||
Example:
|
||||
```python
|
||||
from durabletask import TaskHubGrpcWorker
|
||||
from durabletask.worker import TaskHubGrpcWorker
|
||||
from agent_framework import Agent
|
||||
from agent_framework.azure import DurableAIAgentWorker
|
||||
from agent_framework.azure import AzureOpenAIChatClient
|
||||
from agent_framework_durabletask import DurableAIAgentWorker
|
||||
|
||||
# Create the underlying worker
|
||||
worker = TaskHubGrpcWorker(host_address="localhost:4001")
|
||||
@@ -40,6 +41,7 @@ class DurableAIAgentWorker:
|
||||
agent_worker = DurableAIAgentWorker(worker)
|
||||
|
||||
# Register agents
|
||||
client = AzureOpenAIChatClient()
|
||||
my_agent = Agent(client=client, name="assistant")
|
||||
agent_worker.add_agent(my_agent)
|
||||
|
||||
|
||||
@@ -23,14 +23,14 @@ classifiers = [
|
||||
]
|
||||
dependencies = [
|
||||
"agent-framework-core>=1.0.0rc4",
|
||||
"durabletask>=1.3.0",
|
||||
"durabletask-azuremanaged>=1.3.0",
|
||||
"python-dateutil>=2.8.0",
|
||||
"durabletask>=1.3.0,<2",
|
||||
"durabletask-azuremanaged>=1.3.0,<2",
|
||||
"python-dateutil>=2.8.0,<3",
|
||||
]
|
||||
|
||||
[dependency-groups]
|
||||
dev = [
|
||||
"types-python-dateutil>=2.9.0",
|
||||
"types-python-dateutil==2.9.0.20260305",
|
||||
]
|
||||
|
||||
[tool.uv]
|
||||
@@ -99,7 +99,7 @@ include = "../../shared_tasks.toml"
|
||||
|
||||
[tool.poe.tasks]
|
||||
mypy = "mypy --config-file $POE_ROOT/pyproject.toml agent_framework_durabletask"
|
||||
test = "pytest -m \"not integration\" --cov=agent_framework_durabletask --cov-report=term-missing:skip-covered tests"
|
||||
test = 'pytest -m "not integration" --cov=agent_framework_durabletask --cov-report=term-missing:skip-covered tests'
|
||||
|
||||
[build-system]
|
||||
requires = ["flit-core >= 3.11,<4.0"]
|
||||
|
||||
@@ -2,6 +2,8 @@
|
||||
|
||||
"""Unit tests for AgentSessionId and DurableAgentSession."""
|
||||
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
from agent_framework import AgentSession
|
||||
|
||||
@@ -153,7 +155,7 @@ class TestDurableAgentSession:
|
||||
def test_from_session_id(self) -> None:
|
||||
"""Test creating DurableAgentSession from session ID."""
|
||||
session_id = AgentSessionId(name="TestAgent", key="test-key")
|
||||
session = DurableAgentSession.from_session_id(session_id)
|
||||
session = DurableAgentSession(durable_session_id=session_id)
|
||||
|
||||
assert isinstance(session, DurableAgentSession)
|
||||
assert session.durable_session_id is not None
|
||||
@@ -161,10 +163,10 @@ class TestDurableAgentSession:
|
||||
assert session.durable_session_id.name == "TestAgent"
|
||||
assert session.durable_session_id.key == "test-key"
|
||||
|
||||
def test_from_session_id_with_service_session_id(self) -> None:
|
||||
"""Test creating DurableAgentSession with service session ID."""
|
||||
def test_init_with_service_session_id(self) -> None:
|
||||
"""Test creating DurableAgentSession with explicit service session ID."""
|
||||
session_id = AgentSessionId(name="TestAgent", key="test-key")
|
||||
session = DurableAgentSession.from_session_id(session_id, service_session_id="service-123")
|
||||
session = DurableAgentSession(durable_session_id=session_id, service_session_id="service-123")
|
||||
|
||||
assert session.durable_session_id is not None
|
||||
assert session.durable_session_id == session_id
|
||||
@@ -192,7 +194,7 @@ class TestDurableAgentSession:
|
||||
|
||||
def test_from_dict_with_durable_session_id(self) -> None:
|
||||
"""Test deserialization restores durable session ID."""
|
||||
serialized = {
|
||||
serialized: dict[str, Any] = {
|
||||
"type": "session",
|
||||
"session_id": "session-123",
|
||||
"service_session_id": "service-123",
|
||||
@@ -210,7 +212,7 @@ class TestDurableAgentSession:
|
||||
|
||||
def test_from_dict_without_durable_session_id(self) -> None:
|
||||
"""Test deserialization without durable session ID."""
|
||||
serialized = {
|
||||
serialized: dict[str, Any] = {
|
||||
"type": "session",
|
||||
"session_id": "session-456",
|
||||
"service_session_id": "service-456",
|
||||
|
||||
@@ -88,15 +88,6 @@ class TestDurableAIAgentClientIntegration:
|
||||
|
||||
assert isinstance(session, DurableAgentSession)
|
||||
|
||||
def test_client_agent_session_with_parameters(self, agent_client: DurableAIAgentClient) -> None:
|
||||
"""Verify agent can create sessions with custom parameters."""
|
||||
agent = agent_client.get_agent("assistant")
|
||||
|
||||
session = agent.create_session(service_session_id="client-session-123")
|
||||
|
||||
assert isinstance(session, DurableAgentSession)
|
||||
assert session.service_session_id == "client-session-123"
|
||||
|
||||
|
||||
class TestDurableAIAgentClientPollingConfiguration:
|
||||
"""Test polling configuration parameters for DurableAIAgentClient."""
|
||||
|
||||
@@ -82,17 +82,6 @@ class TestDurableAIAgentOrchestrationContextIntegration:
|
||||
|
||||
assert isinstance(session, DurableAgentSession)
|
||||
|
||||
def test_orchestration_agent_session_with_parameters(
|
||||
self, agent_context: DurableAIAgentOrchestrationContext
|
||||
) -> None:
|
||||
"""Verify agent can create sessions with custom parameters."""
|
||||
agent = agent_context.get_agent("assistant")
|
||||
|
||||
session = agent.create_session(service_session_id="orch-session-456")
|
||||
|
||||
assert isinstance(session, DurableAgentSession)
|
||||
assert session.service_session_id == "orch-session-456"
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
pytest.main([__file__, "-v", "--tb=short"])
|
||||
|
||||
@@ -184,16 +184,31 @@ class TestDurableAIAgentSessionManagement:
|
||||
mock_executor.get_new_session.assert_called_once_with("test_agent")
|
||||
assert session == mock_session
|
||||
|
||||
def test_create_session_forwards_kwargs(self, test_agent: DurableAIAgent[Any], mock_executor: Mock) -> None:
|
||||
"""Verify create_session forwards kwargs to executor."""
|
||||
mock_session = DurableAgentSession(service_session_id="session-123")
|
||||
def test_get_session_forwards_service_session_id(
|
||||
self, test_agent: DurableAIAgent[Any], mock_executor: Mock
|
||||
) -> None:
|
||||
"""Verify get_session forwards service_session_id and session_id to executor."""
|
||||
mock_session = DurableAgentSession(service_session_id="svc-123")
|
||||
mock_executor.get_new_session.return_value = mock_session
|
||||
|
||||
test_agent.create_session(service_session_id="session-123")
|
||||
session = test_agent.get_session("svc-123", session_id="local-456")
|
||||
|
||||
mock_executor.get_new_session.assert_called_once()
|
||||
_, kwargs = mock_executor.get_new_session.call_args
|
||||
assert kwargs["service_session_id"] == "session-123"
|
||||
mock_executor.get_new_session.assert_called_once_with(
|
||||
"test_agent", service_session_id="svc-123", session_id="local-456"
|
||||
)
|
||||
assert session.service_session_id == "svc-123"
|
||||
|
||||
def test_get_session_without_session_id(self, test_agent: DurableAIAgent[Any], mock_executor: Mock) -> None:
|
||||
"""Verify get_session works with only service_session_id (session_id defaults to None)."""
|
||||
mock_session = DurableAgentSession(service_session_id="svc-789")
|
||||
mock_executor.get_new_session.return_value = mock_session
|
||||
|
||||
session = test_agent.get_session("svc-789")
|
||||
|
||||
mock_executor.get_new_session.assert_called_once_with(
|
||||
"test_agent", service_session_id="svc-789", session_id=None
|
||||
)
|
||||
assert session.service_session_id == "svc-789"
|
||||
|
||||
|
||||
class TestDurableAgentProviderInterface:
|
||||
|
||||
+3
-4
@@ -146,11 +146,11 @@ class FoundryLocalClient(
|
||||
timeout: float | None = None,
|
||||
prepare_model: bool = True,
|
||||
device: DeviceType | None = None,
|
||||
additional_properties: dict[str, Any] | None = None,
|
||||
middleware: Sequence[ChatAndFunctionMiddlewareTypes] | None = None,
|
||||
function_invocation_configuration: FunctionInvocationConfiguration | None = None,
|
||||
env_file_path: str | None = None,
|
||||
env_file_encoding: str = "utf-8",
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
"""Initialize a FoundryLocalClient.
|
||||
|
||||
@@ -169,12 +169,11 @@ class FoundryLocalClient(
|
||||
The device is used to select the appropriate model variant.
|
||||
If not provided, the default device for your system will be used.
|
||||
The values are in the foundry_local.models.DeviceType enum.
|
||||
additional_properties: Additional properties stored on the client instance.
|
||||
middleware: Optional sequence of ChatAndFunctionMiddlewareTypes to apply to requests.
|
||||
function_invocation_configuration: Optional configuration for function invocation support.
|
||||
env_file_path: If provided, the .env settings are read from this file path location.
|
||||
env_file_encoding: The encoding of the .env file, defaults to 'utf-8'.
|
||||
kwargs: Additional keyword arguments, are passed to the RawOpenAIChatClient.
|
||||
This can include middleware and additional properties.
|
||||
|
||||
Examples:
|
||||
|
||||
@@ -271,8 +270,8 @@ class FoundryLocalClient(
|
||||
super().__init__(
|
||||
model_id=model_info.id,
|
||||
client=AsyncOpenAI(base_url=manager.endpoint, api_key=manager.api_key),
|
||||
additional_properties=additional_properties,
|
||||
middleware=middleware,
|
||||
function_invocation_configuration=function_invocation_configuration,
|
||||
**kwargs,
|
||||
)
|
||||
self.manager = manager
|
||||
|
||||
@@ -24,7 +24,7 @@ classifiers = [
|
||||
]
|
||||
dependencies = [
|
||||
"agent-framework-core>=1.0.0rc4",
|
||||
"foundry-local-sdk>=0.5.1,<1",
|
||||
"foundry-local-sdk>=0.5.1,<0.5.2",
|
||||
]
|
||||
|
||||
[tool.uv]
|
||||
@@ -86,7 +86,7 @@ include = "../../shared_tasks.toml"
|
||||
|
||||
[tool.poe.tasks]
|
||||
mypy = "mypy --config-file $POE_ROOT/pyproject.toml agent_framework_foundry_local"
|
||||
test = "pytest -m \"not integration\" --cov=agent_framework_foundry_local --cov-report=term-missing:skip-covered tests"
|
||||
test = 'pytest -m "not integration" --cov=agent_framework_foundry_local --cov-report=term-missing:skip-covered tests'
|
||||
|
||||
[build-system]
|
||||
requires = ["flit-core >= 3.11,<4.0"]
|
||||
|
||||
@@ -25,20 +25,27 @@ from agent_framework._settings import load_settings
|
||||
from agent_framework._tools import FunctionTool, ToolTypes
|
||||
from agent_framework._types import AgentRunInputs, normalize_tools
|
||||
from agent_framework.exceptions import AgentException
|
||||
from copilot import CopilotClient, CopilotSession
|
||||
from copilot.generated.session_events import PermissionRequest, SessionEvent, SessionEventType
|
||||
from copilot.types import (
|
||||
CopilotClientOptions,
|
||||
MCPServerConfig,
|
||||
MessageOptions,
|
||||
PermissionRequestResult,
|
||||
ResumeSessionConfig,
|
||||
SessionConfig,
|
||||
SystemMessageConfig,
|
||||
ToolInvocation,
|
||||
ToolResult,
|
||||
)
|
||||
from copilot.types import Tool as CopilotTool
|
||||
|
||||
try:
|
||||
from copilot import CopilotClient, CopilotSession
|
||||
from copilot.generated.session_events import PermissionRequest, SessionEvent, SessionEventType
|
||||
from copilot.types import (
|
||||
CopilotClientOptions,
|
||||
MCPServerConfig,
|
||||
MessageOptions,
|
||||
PermissionRequestResult,
|
||||
ResumeSessionConfig,
|
||||
SessionConfig,
|
||||
SystemMessageConfig,
|
||||
ToolInvocation,
|
||||
ToolResult,
|
||||
)
|
||||
from copilot.types import Tool as CopilotTool
|
||||
except ImportError as _copilot_import_error:
|
||||
raise ImportError(
|
||||
"GitHubCopilotAgent requires the 'github-copilot-sdk' package, which is only available on Python 3.11+. "
|
||||
"Please use Python 3.11 or later."
|
||||
) from _copilot_import_error
|
||||
|
||||
if sys.version_info >= (3, 13):
|
||||
from typing import TypeVar
|
||||
@@ -303,7 +310,6 @@ class GitHubCopilotAgent(BaseAgent, Generic[OptionsT]):
|
||||
stream: Literal[False] = False,
|
||||
session: AgentSession | None = None,
|
||||
options: OptionsT | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Awaitable[AgentResponse]: ...
|
||||
|
||||
@overload
|
||||
@@ -314,7 +320,6 @@ class GitHubCopilotAgent(BaseAgent, Generic[OptionsT]):
|
||||
stream: Literal[True],
|
||||
session: AgentSession | None = None,
|
||||
options: OptionsT | None = None,
|
||||
**kwargs: Any,
|
||||
) -> ResponseStream[AgentResponseUpdate, AgentResponse]: ...
|
||||
|
||||
def run(
|
||||
@@ -324,7 +329,6 @@ class GitHubCopilotAgent(BaseAgent, Generic[OptionsT]):
|
||||
stream: bool = False,
|
||||
session: AgentSession | None = None,
|
||||
options: OptionsT | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Awaitable[AgentResponse] | ResponseStream[AgentResponseUpdate, AgentResponse]:
|
||||
"""Get a response from the agent.
|
||||
|
||||
@@ -339,7 +343,6 @@ class GitHubCopilotAgent(BaseAgent, Generic[OptionsT]):
|
||||
stream: Whether to stream the response. Defaults to False.
|
||||
session: The conversation session associated with the message(s).
|
||||
options: Runtime options (model, timeout, etc.).
|
||||
kwargs: Additional keyword arguments.
|
||||
|
||||
Returns:
|
||||
When stream=False: An Awaitable[AgentResponse].
|
||||
@@ -354,10 +357,10 @@ class GitHubCopilotAgent(BaseAgent, Generic[OptionsT]):
|
||||
return AgentResponse.from_updates(updates)
|
||||
|
||||
return ResponseStream(
|
||||
self._stream_updates(messages=messages, session=session, options=options, **kwargs),
|
||||
self._stream_updates(messages=messages, session=session, options=options),
|
||||
finalizer=_finalize,
|
||||
)
|
||||
return self._run_impl(messages=messages, session=session, options=options, **kwargs)
|
||||
return self._run_impl(messages=messages, session=session, options=options)
|
||||
|
||||
async def _run_impl(
|
||||
self,
|
||||
@@ -365,7 +368,6 @@ class GitHubCopilotAgent(BaseAgent, Generic[OptionsT]):
|
||||
*,
|
||||
session: AgentSession | None = None,
|
||||
options: OptionsT | None = None,
|
||||
**kwargs: Any,
|
||||
) -> AgentResponse:
|
||||
"""Non-streaming implementation of run."""
|
||||
if not self._started:
|
||||
@@ -414,7 +416,6 @@ class GitHubCopilotAgent(BaseAgent, Generic[OptionsT]):
|
||||
*,
|
||||
session: AgentSession | None = None,
|
||||
options: OptionsT | None = None,
|
||||
**kwargs: Any,
|
||||
) -> AsyncIterable[AgentResponseUpdate]:
|
||||
"""Internal method to stream updates from GitHub Copilot.
|
||||
|
||||
@@ -424,7 +425,6 @@ class GitHubCopilotAgent(BaseAgent, Generic[OptionsT]):
|
||||
Keyword Args:
|
||||
session: The conversation session associated with the message(s).
|
||||
options: Runtime options (model, timeout, etc.).
|
||||
kwargs: Additional keyword arguments.
|
||||
|
||||
Yields:
|
||||
AgentResponseUpdate items.
|
||||
|
||||
@@ -3,7 +3,7 @@ name = "agent-framework-github-copilot"
|
||||
description = "GitHub Copilot integration for Microsoft Agent Framework."
|
||||
authors = [{ name = "Microsoft", email = "af-support@microsoft.com"}]
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.11"
|
||||
requires-python = ">=3.10"
|
||||
version = "1.0.0b260311"
|
||||
license-files = ["LICENSE"]
|
||||
urls.homepage = "https://aka.ms/agent-framework"
|
||||
@@ -15,6 +15,7 @@ classifiers = [
|
||||
"Development Status :: 4 - Beta",
|
||||
"Intended Audience :: Developers",
|
||||
"Programming Language :: Python :: 3",
|
||||
"Programming Language :: Python :: 3.10",
|
||||
"Programming Language :: Python :: 3.11",
|
||||
"Programming Language :: Python :: 3.12",
|
||||
"Programming Language :: Python :: 3.13",
|
||||
@@ -23,7 +24,7 @@ classifiers = [
|
||||
]
|
||||
dependencies = [
|
||||
"agent-framework-core>=1.0.0rc4",
|
||||
"github-copilot-sdk>=0.1.32",
|
||||
"github-copilot-sdk>=0.1.31,<0.1.33; python_version >= '3.11'",
|
||||
]
|
||||
|
||||
[tool.uv]
|
||||
@@ -85,9 +86,16 @@ executor.type = "uv"
|
||||
include = "../../shared_tasks.toml"
|
||||
|
||||
[tool.poe.tasks]
|
||||
mypy = "mypy --config-file $POE_ROOT/pyproject.toml agent_framework_github_copilot"
|
||||
test = "pytest -m \"not integration\" --cov=agent_framework_github_copilot --cov-report=term-missing:skip-covered tests"
|
||||
|
||||
[tool.poe.tasks.pyright]
|
||||
shell = "python -c \"import sys; exit(0 if sys.version_info < (3,11) else 1)\" || pyright"
|
||||
interpreter = "posix"
|
||||
|
||||
[tool.poe.tasks.mypy]
|
||||
shell = "python -c \"import sys; exit(0 if sys.version_info < (3,11) else 1)\" || mypy --config-file $POE_ROOT/pyproject.toml agent_framework_github_copilot"
|
||||
interpreter = "posix"
|
||||
|
||||
[build-system]
|
||||
requires = ["flit-core >= 3.11,<4.0"]
|
||||
build-backend = "flit_core.buildapi"
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
# ruff: noqa: E402
|
||||
|
||||
import unittest.mock
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any
|
||||
@@ -7,6 +9,9 @@ from unittest.mock import AsyncMock, MagicMock, patch
|
||||
from uuid import uuid4
|
||||
|
||||
import pytest
|
||||
|
||||
copilot = pytest.importorskip("copilot")
|
||||
|
||||
from agent_framework import (
|
||||
AgentResponse,
|
||||
AgentResponseUpdate,
|
||||
|
||||
@@ -62,6 +62,27 @@ For example, to use the GAIA module:
|
||||
from agent_framework.lab.gaia import GAIA
|
||||
```
|
||||
|
||||
## Running Tests Locally
|
||||
|
||||
For machine-safe local runs, prefer package-scoped commands first:
|
||||
|
||||
```bash
|
||||
uv run --directory packages/lab poe test
|
||||
uv run --directory packages/lab pytest -q -m "not integration"
|
||||
```
|
||||
|
||||
When you need to run package tasks from the repository root, use sequential mode to avoid launching all package tests in parallel:
|
||||
|
||||
```bash
|
||||
uv run poe test --seq
|
||||
```
|
||||
|
||||
Lightning observability tests intentionally exercise heavier tracing paths and are marked as `resource_intensive`:
|
||||
|
||||
```bash
|
||||
uv run --directory packages/lab pytest lightning/tests/test_lightning.py -m "resource_intensive" -q
|
||||
```
|
||||
|
||||
## Should I consume Lab Modules?
|
||||
|
||||
If you are looking for stable and production-ready features, you should not use lab modules. Stick to the core framework.
|
||||
|
||||
@@ -10,10 +10,11 @@ import re
|
||||
import string
|
||||
import tempfile
|
||||
import time
|
||||
from collections.abc import Iterable
|
||||
from collections.abc import Callable, Iterable
|
||||
from datetime import datetime
|
||||
from functools import lru_cache
|
||||
from pathlib import Path
|
||||
from typing import Any, cast
|
||||
from typing import Any, Protocol, cast
|
||||
|
||||
from opentelemetry.trace import NoOpTracer, SpanKind, get_tracer
|
||||
from tqdm import tqdm
|
||||
@@ -23,6 +24,33 @@ from ._types import Evaluation, Evaluator, Prediction, Task, TaskResult, TaskRun
|
||||
__all__ = ["GAIA", "GAIATelemetryConfig", "gaia_scorer"]
|
||||
|
||||
|
||||
class _OrjsonModule(Protocol):
|
||||
def dumps(self, obj: object, /, default: Callable[[Any], object] | None = None) -> bytes: ...
|
||||
|
||||
def loads(self, obj: str | bytes | bytearray, /) -> object: ...
|
||||
|
||||
|
||||
@lru_cache(maxsize=1)
|
||||
def _get_orjson() -> _OrjsonModule | None:
|
||||
try:
|
||||
import orjson as runtime_orjson # pyright: ignore[reportMissingImports]
|
||||
except ImportError:
|
||||
return None
|
||||
return cast(_OrjsonModule, runtime_orjson)
|
||||
|
||||
|
||||
def _dump_json_line(value: object) -> str:
|
||||
if (runtime_orjson := _get_orjson()) is not None:
|
||||
return runtime_orjson.dumps(value, default=str).decode("utf-8")
|
||||
return json.dumps(value, default=str)
|
||||
|
||||
|
||||
def _load_json_value(value: str | bytes) -> object:
|
||||
if (runtime_orjson := _get_orjson()) is not None:
|
||||
return runtime_orjson.loads(value)
|
||||
return json.loads(value)
|
||||
|
||||
|
||||
class GAIATelemetryConfig:
|
||||
"""Configuration for GAIA telemetry and tracing."""
|
||||
|
||||
@@ -226,13 +254,7 @@ def _read_jsonl(path: Path) -> Iterable[dict[str, Any]]:
|
||||
for line in f:
|
||||
if not line.strip():
|
||||
continue
|
||||
parsed: object
|
||||
try:
|
||||
import orjson
|
||||
|
||||
parsed = orjson.loads(line)
|
||||
except Exception:
|
||||
parsed = json.loads(line)
|
||||
parsed = _load_json_value(line)
|
||||
|
||||
record = _coerce_record(parsed)
|
||||
if record is not None:
|
||||
@@ -620,12 +642,7 @@ class GAIA:
|
||||
"prediction_metadata": result.prediction.metadata,
|
||||
"evaluation_details": result.evaluation.details,
|
||||
}
|
||||
try:
|
||||
import orjson
|
||||
|
||||
f.write(orjson.dumps(record, default=str).decode("utf-8") + "\n")
|
||||
except ImportError:
|
||||
f.write(json.dumps(record, default=str) + "\n")
|
||||
f.write(_dump_json_line(record) + "\n")
|
||||
|
||||
|
||||
def viewer_main() -> None:
|
||||
@@ -646,13 +663,7 @@ def viewer_main() -> None:
|
||||
with open(args.results_file, encoding="utf-8") as f:
|
||||
for line in f:
|
||||
if line.strip():
|
||||
try:
|
||||
import orjson
|
||||
|
||||
parsed: object = orjson.loads(line)
|
||||
except ImportError:
|
||||
parsed = json.loads(line)
|
||||
|
||||
parsed = _load_json_value(line)
|
||||
record = _coerce_record(parsed)
|
||||
if record is not None:
|
||||
results.append(record)
|
||||
|
||||
@@ -2,10 +2,14 @@
|
||||
|
||||
"""RL Module for Microsoft Agent Framework."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import importlib.metadata
|
||||
|
||||
from agent_framework.observability import enable_instrumentation
|
||||
from agentlightning import AgentOpsTracer # type: ignore
|
||||
from agentlightning.tracer import (
|
||||
AgentOpsTracer, # pyright: ignore[reportMissingImports] # type: ignore[import-not-found]
|
||||
)
|
||||
|
||||
try:
|
||||
__version__ = importlib.metadata.version(__name__)
|
||||
@@ -23,11 +27,11 @@ class AgentFrameworkTracer(AgentOpsTracer): # type: ignore
|
||||
def init(self) -> None:
|
||||
"""Initialize the agent-framework-lab-lightning for training."""
|
||||
enable_instrumentation()
|
||||
super().init()
|
||||
super().init() # pyright: ignore[reportUnknownMemberType]
|
||||
|
||||
def teardown(self) -> None:
|
||||
"""Teardown the agent-framework-lab-lightning for training."""
|
||||
super().teardown()
|
||||
super().teardown() # pyright: ignore[reportUnknownMemberType]
|
||||
|
||||
|
||||
__all__: list[str] = ["AgentFrameworkTracer"]
|
||||
|
||||
@@ -7,12 +7,8 @@ from unittest.mock import AsyncMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
agentlightning = pytest.importorskip("agentlightning")
|
||||
|
||||
from agent_framework import AgentExecutor, AgentResponse, Agent, WorkflowBuilder, Workflow
|
||||
from agent_framework_lab_lightning import AgentFrameworkTracer
|
||||
from agent_framework.openai import OpenAIChatClient
|
||||
from agentlightning import TracerTraceToTriplet
|
||||
from openai.types.chat import ChatCompletion, ChatCompletionMessage
|
||||
from openai.types.chat.chat_completion import Choice
|
||||
|
||||
@@ -118,6 +114,7 @@ async def test_openai_workflow_two_agents(workflow_two_agents: Workflow):
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.resource_intensive
|
||||
async def test_observability(workflow_two_agents: Workflow):
|
||||
r"""Expected trace tree:
|
||||
|
||||
@@ -129,6 +126,10 @@ async def test_observability(workflow_two_agents: Workflow):
|
||||
| |
|
||||
[chat gpt-4o] [chat gpt-4o]
|
||||
"""
|
||||
pytest.importorskip("agentlightning")
|
||||
from agent_framework_lab_lightning import AgentFrameworkTracer
|
||||
from agentlightning.adapter import TracerTraceToTriplet
|
||||
|
||||
tracer = AgentFrameworkTracer()
|
||||
try:
|
||||
tracer.init()
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user