Compare commits

...
Author SHA1 Message Date
Shawn HenryandGitHub 23c8d97f21 New Microsoft Agent Framework logos 2026-06-07 20:40:13 -07:00
Yufeng HeandGitHub fa9e086576 fix: preserve foreach record values (#6208) 2026-06-05 22:01:59 +00:00
dcc218dbac Python: feat(python): Add MCP client OTel spans per GenAI semantic conventions (#6349)
* feat(python): Add MCP client OTel spans per GenAI semantic conventions

Implement MCP client spans per the OTel GenAI Semantic Conventions for MCP
(https://opentelemetry.io/docs/specs/semconv/gen-ai/mcp/#client).

Operations instrumented:
- initialize: CLIENT span capturing MCP session setup
- tools/list: CLIENT span for tool listing (per-page)
- prompts/list: CLIENT span for prompt listing (per-page)
- tools/call: CLIENT span (nested under execute_tool when called via FunctionTool)
- prompts/get: CLIENT span

Span attributes follow the MCP semantic conventions:
- Required: mcp.method.name
- Conditional: error.type, gen_ai.tool.name, gen_ai.prompt.name
- Recommended: gen_ai.operation.name, mcp.protocol.version, mcp.session.id,
  network.transport, server.address, server.port

Transport-specific attributes per subclass:
- MCPStdioTool: network.transport=pipe
- MCPStreamableHTTPTool: network.transport=tcp, network.protocol.name=http
- MCPWebsocketTool: network.transport=tcp, network.protocol.name=websocket

All span creation gated behind OBSERVABILITY_SETTINGS.ENABLED.

Closes #3624
Closes #4697

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>

* refactor: simplify MCP spans — remove enrichment logic and protocol version caching

- Always create nested CLIENT spans for tools/call instead of enriching
  the parent execute_tool span
- Remove _ACTIVE_TOOL_EXECUTION_SPAN contextvar (no longer needed)
- Remove enrich_span_with_mcp_attributes() helper
- Remove _otel_error_type preservation in FunctionTool.invoke()
- Remove _mcp_protocol_version instance variable; protocol version is
  only set on the initialize span where it is available

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>

* Refine copilot solution

* fix: enable automatic exception recording on MCP spans

Remove record_exception=False and set_status_on_exception=False from
create_mcp_client_span. Let OTel handle exception recording and status
setting automatically. The manual set_mcp_span_error calls for tools/call
still correctly set error.type (which OTel's automatic handling doesn't
touch), so tool_error is preserved.

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>

* Reduce number of lines

* Add comment to sample

* test: address PR review comments on MCP observability tests

- Fix initialize test to call mocked session.initialize() and read
  protocolVersion from the result instead of hardcoding it
- Add tools/call McpError error-path test
- Add prompts/get McpError error-path test

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>

* Fix export error

---------

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
2026-06-05 19:23:01 +00:00
6bd2cfec03 .NET: [BREAKING] Add auto-approval rules (heuristics) to ToolApprovalAgent (#6335)
* Add support for approving tools via heuristic rules

* Address PR comments

* Address PR comments

* Apply suggestion from @SergeyMenshykh

Co-authored-by: SergeyMenshykh <68852919+SergeyMenshykh@users.noreply.github.com>

---------

Co-authored-by: SergeyMenshykh <68852919+SergeyMenshykh@users.noreply.github.com>
2026-06-05 18:43:07 +01:00
westeyandGitHub ab8ba8fc61 .NET: Allow storage of auto-approved functions (#4950)
* Allow storage of auto-approved functions

* Address PR comments
2026-06-05 18:42:21 +01:00
Tao ChenandGitHub 9cafd7e58b Python: Refactor workflow as agent pending request handling (#6259)
* WIP: Refactor Workflow as agent pending request handling

* WIP: debugging empty message bug

* Working: Workflow as agent with function approval

* Address Copilot comments

* Fix mypy

* Address comments and fix pipeline

* Request info non function approval now becomes function call

* Revert uv.lock

* Fix mypy

* Bump min version of azure-ai-project

* Remove RequestInfoFunctionArgs

* fix tests

* Fix failing tests

* Fix sample
2026-06-05 17:23:19 +00:00
d5335fbeae Python (fix:gemini): make Gemini honor declarative outputSchema, not just JSON mode (#5893)
* fix(gemini): preserve schema response_format

* fix(gemini): satisfy pyright strict in response schema extraction

Cast Any-narrowed mappings to Mapping[str, Any] in the structured-output
schema helpers so pyright strict no longer reports partially-unknown
member, argument, and variable types. Pass response_format["format"]
straight into the recursive extractor, which already guards non-mapping
inputs. No behavior change.

* fix(gemini): use Sequence[object] cast to satisfy both mypy and pyright

The Sequence[Any] cast pyright strict needs to know the loop element type
is reported as a redundant-cast by mypy, which already narrows the
isinstance branch to Sequence[Any]. Cast to Sequence[object] instead:
pyright gets a fully known element type and mypy no longer sees an
identical-type cast. No behavior change.

---------

Co-authored-by: Evan Mattson <evan.mattson@microsoft.com>
2026-06-05 15:17:51 +00:00
bf4ad48cf2 Python: MCP long-running task support in Python (#6319)
* MCP long-running task support in Python

* Fix pyupgrade and AGENTS.md reconnect description

- pyupgrade: drop forward-reference string annotations in _mcp.py (Python 3.10+ resolves them natively now that MCPTaskOptions is defined before use).

- AGENTS.md: align reconnect description with current behavior. Phase 1 (initial tools/call) does NOT retry on connection loss; raises 'connection lost; task state unknown' instead, so a server that accepted the request but lost the response cannot start the operation twice. Phase 2 (tasks/get / tasks/result) still reconnects once against the same task_id.

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>

* Fix bandit nosec marker for CI pipeline

* Address PR feedbacks

* Clarifiied comments and addressed more PR feedbacks.

---------

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
2026-06-05 00:04:55 +00:00
52 changed files with 6211 additions and 432 deletions
Binary file not shown.

After

Width:  |  Height:  |  Size: 219 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 15 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 16 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 19 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 20 KiB

@@ -0,0 +1,55 @@
<svg width="256" height="256" viewBox="0 0 256 256" fill="none" xmlns="http://www.w3.org/2000/svg">
<rect width="256" height="256" transform="matrix(-1 0 0 1 256 0)" fill="white"/>
<path d="M163.006 36.0262C155.129 31.4785 145.424 31.4759 137.533 36.0284L104.002 55.3877C107.838 53.2763 112.07 52.2274 116.296 52.2248C120.633 52.2231 124.976 53.3247 128.871 55.5411L174.091 81.6479C175.11 82.2362 175.738 83.3236 175.738 84.5004V151.263C175.738 152.716 177.313 153.621 178.568 152.895L196.967 142.28C204.845 137.732 209.704 129.326 209.704 120.221V77.7299C209.704 68.6342 204.855 60.227 196.967 55.6697L190.983 52.2141C190.65 52.0008 190.313 51.7921 189.969 51.5933L163.006 36.0262Z" fill="url(#paint0_linear_481_4810)"/>
<path d="M163.006 36.0262C155.129 31.4785 145.424 31.4759 137.533 36.0284L104.002 55.3877C107.838 53.2763 112.07 52.2274 116.296 52.2248C120.633 52.2231 124.976 53.3247 128.871 55.5411L174.091 81.6479C175.11 82.2362 175.738 83.3236 175.738 84.5004V151.263C175.738 152.716 177.313 153.621 178.568 152.895L196.967 142.28C204.845 137.732 209.704 129.326 209.704 120.221V77.7299C209.704 68.6342 204.855 60.227 196.967 55.6697L190.983 52.2141C190.65 52.0008 190.313 51.7921 189.969 51.5933L163.006 36.0262Z" fill="url(#paint1_linear_481_4810)"/>
<path d="M103.548 55.6397L103.557 55.6451L104.002 55.3877C103.851 55.471 103.698 55.5531 103.548 55.6397Z" fill="url(#paint2_linear_481_4810)"/>
<path d="M103.548 55.6397L103.557 55.6451L104.002 55.3877C103.851 55.471 103.698 55.5531 103.548 55.6397Z" fill="url(#paint3_linear_481_4810)"/>
<path d="M116.308 52.2209C111.903 52.2271 107.507 53.3498 103.561 55.6366C95.6702 60.1891 90.8239 68.6019 90.8231 77.6986L90.8223 167.846C90.8222 169.786 92.871 171.041 94.599 170.16L103.523 165.611C116.573 158.959 124.788 145.549 124.788 130.902V62.9894C124.778 58.6476 129.49 55.9209 133.25 58.0698L128.879 55.5453C124.976 53.3242 120.645 52.2192 116.308 52.2209Z" fill="url(#paint4_linear_481_4810)"/>
<path d="M95.3068 221.682C103.184 226.23 112.889 226.232 120.78 221.68L154.311 202.32C150.475 204.432 146.243 205.481 142.018 205.483C137.68 205.485 133.337 204.383 129.442 202.167L84.2226 176.06C83.2035 175.472 82.5757 174.384 82.5757 173.208L82.5757 106.445C82.5757 104.992 81 104.087 79.7451 104.813L61.3465 115.428C53.4682 119.976 48.6091 128.382 48.6089 137.487L48.6089 179.978C48.6089 189.074 53.4585 197.481 61.3465 202.038L67.3303 205.494C67.6629 205.707 68.0002 205.916 68.3446 206.115L95.3068 221.682Z" fill="url(#paint5_linear_481_4810)"/>
<path d="M95.3068 221.682C103.184 226.23 112.889 226.232 120.78 221.68L154.311 202.32C150.475 204.432 146.243 205.481 142.018 205.483C137.68 205.485 133.337 204.383 129.442 202.167L84.2226 176.06C83.2035 175.472 82.5757 174.384 82.5757 173.208L82.5757 106.445C82.5757 104.992 81 104.087 79.7451 104.813L61.3465 115.428C53.4682 119.976 48.6091 128.382 48.6089 137.487L48.6089 179.978C48.6089 189.074 53.4585 197.481 61.3465 202.038L67.3303 205.494C67.6629 205.707 68.0002 205.916 68.3446 206.115L95.3068 221.682Z" fill="url(#paint6_linear_481_4810)"/>
<path d="M154.765 202.068L154.756 202.063L154.311 202.32C154.463 202.237 154.615 202.155 154.765 202.068Z" fill="url(#paint7_linear_481_4810)"/>
<path d="M154.765 202.068L154.756 202.063L154.311 202.32C154.463 202.237 154.615 202.155 154.765 202.068Z" fill="url(#paint8_linear_481_4810)"/>
<path d="M142.003 205.487C146.408 205.481 150.805 204.358 154.751 202.071C162.641 197.519 167.488 189.106 167.488 180.009L167.489 89.8618C167.489 87.9222 165.44 86.667 163.712 87.5479L154.788 92.0972C141.739 98.7494 133.523 112.159 133.523 126.806L133.523 194.719C133.533 199.06 128.821 201.787 125.061 199.638L129.432 202.163C133.336 204.384 137.666 205.489 142.003 205.487Z" fill="url(#paint9_linear_481_4810)"/>
<defs>
<linearGradient id="paint0_linear_481_4810" x1="207.413" y1="61.882" x2="148.999" y2="163.067" gradientUnits="userSpaceOnUse">
<stop stop-color="#9189F7"/>
<stop offset="1" stop-color="#4135E9"/>
</linearGradient>
<linearGradient id="paint1_linear_481_4810" x1="188.235" y1="204.189" x2="264.21" y2="152.592" gradientUnits="userSpaceOnUse">
<stop stop-color="#4F42FD"/>
<stop offset="1" stop-color="#7274FF"/>
</linearGradient>
<linearGradient id="paint2_linear_481_4810" x1="207.413" y1="61.882" x2="148.999" y2="163.067" gradientUnits="userSpaceOnUse">
<stop stop-color="#9189F7"/>
<stop offset="1" stop-color="#4135E9"/>
</linearGradient>
<linearGradient id="paint3_linear_481_4810" x1="188.235" y1="204.189" x2="264.21" y2="152.592" gradientUnits="userSpaceOnUse">
<stop stop-color="#4F42FD"/>
<stop offset="1" stop-color="#7274FF"/>
</linearGradient>
<linearGradient id="paint4_linear_481_4810" x1="93.1761" y1="128.826" x2="66.2399" y2="104.746" gradientUnits="userSpaceOnUse">
<stop offset="0.25" stop-color="#4F42FD"/>
<stop offset="1" stop-color="#2C08AC"/>
</linearGradient>
<linearGradient id="paint5_linear_481_4810" x1="50.9004" y1="195.826" x2="109.315" y2="94.6412" gradientUnits="userSpaceOnUse">
<stop stop-color="#9189F7"/>
<stop offset="1" stop-color="#4135E9"/>
</linearGradient>
<linearGradient id="paint6_linear_481_4810" x1="70.0787" y1="53.5188" x2="-5.89657" y2="105.116" gradientUnits="userSpaceOnUse">
<stop stop-color="#4F42FD"/>
<stop offset="1" stop-color="#7274FF"/>
</linearGradient>
<linearGradient id="paint7_linear_481_4810" x1="50.9004" y1="195.826" x2="109.315" y2="94.6412" gradientUnits="userSpaceOnUse">
<stop stop-color="#9189F7"/>
<stop offset="1" stop-color="#4135E9"/>
</linearGradient>
<linearGradient id="paint8_linear_481_4810" x1="70.0787" y1="53.5188" x2="-5.89657" y2="105.116" gradientUnits="userSpaceOnUse">
<stop stop-color="#4F42FD"/>
<stop offset="1" stop-color="#7274FF"/>
</linearGradient>
<linearGradient id="paint9_linear_481_4810" x1="165.135" y1="128.882" x2="192.072" y2="152.962" gradientUnits="userSpaceOnUse">
<stop offset="0.25" stop-color="#4F42FD"/>
<stop offset="1" stop-color="#2C08AC"/>
</linearGradient>
</defs>
</svg>

After

Width:  |  Height:  |  Size: 5.8 KiB

@@ -0,0 +1,5 @@
<svg width="256" height="256" viewBox="0 0 256 256" fill="none" xmlns="http://www.w3.org/2000/svg">
<rect width="256" height="256" transform="matrix(-1 0 0 1 256 0)" fill="white"/>
<path d="M167.489 89.8618C167.489 87.9224 165.441 86.667 163.713 87.5473L154.788 92.0971C141.739 98.7493 133.523 112.159 133.523 126.806V194.718C133.533 199.053 128.837 201.777 125.08 199.647L84.2227 176.06C83.2036 175.472 82.5762 174.384 82.5762 173.207V106.445C82.5759 104.992 80.9999 104.087 79.7451 104.813L61.3467 115.428C53.4685 119.976 48.6096 128.382 48.6094 137.487V179.978C48.6094 189.074 53.4588 197.481 61.3467 202.039L67.3301 205.494C67.6627 205.707 68.0004 205.916 68.3447 206.115L95.3066 221.682C103.184 226.23 112.889 226.232 120.779 221.679L154.312 202.32C154.271 202.342 154.23 202.362 154.189 202.384C154.283 202.333 154.375 202.281 154.468 202.229L154.312 202.32C154.463 202.237 154.615 202.154 154.765 202.068L154.762 202.065C162.646 197.511 167.487 189.102 167.488 180.009L167.489 89.8618Z" fill="black"/>
<path d="M163.007 36.0259C155.13 31.4781 145.424 31.4763 137.533 36.0288L104.002 55.3882C104.029 55.3732 104.057 55.359 104.084 55.3442C104.02 55.3794 103.956 55.4159 103.893 55.4516L104.002 55.3882C103.851 55.4714 103.699 55.5536 103.549 55.6401L103.552 55.6411C95.6664 60.1948 90.8241 68.6053 90.8232 77.6987L90.8223 167.846C90.8223 169.786 92.8707 171.04 94.5986 170.16L103.523 165.611C116.573 158.959 124.788 145.549 124.788 130.902V62.9897C124.778 58.6479 129.49 55.9211 133.25 58.0698L174.091 81.6479C175.11 82.2363 175.737 83.3238 175.737 84.5005V151.263C175.738 152.716 177.314 153.621 178.568 152.895L196.967 142.28C204.845 137.732 209.704 129.326 209.704 120.221V77.73C209.704 68.6343 204.855 60.2267 196.967 55.6694L190.982 52.2143C190.65 52.0011 190.313 51.792 189.969 51.5932L163.007 36.0259Z" fill="black"/>
</svg>

After

Width:  |  Height:  |  Size: 1.8 KiB

@@ -0,0 +1,5 @@
<svg width="256" height="256" viewBox="0 0 256 256" fill="none" xmlns="http://www.w3.org/2000/svg">
<rect width="256" height="256" transform="matrix(-1 0 0 1 256 0)" fill="black"/>
<path d="M167.489 89.8618C167.489 87.9224 165.441 86.667 163.713 87.5473L154.788 92.0971C141.739 98.7493 133.523 112.159 133.523 126.806V194.718C133.533 199.053 128.837 201.777 125.08 199.647L84.2227 176.06C83.2036 175.472 82.5762 174.384 82.5762 173.207V106.445C82.5759 104.992 80.9999 104.087 79.7451 104.813L61.3467 115.428C53.4685 119.976 48.6096 128.382 48.6094 137.487V179.978C48.6094 189.074 53.4588 197.481 61.3467 202.039L67.3301 205.494C67.6627 205.707 68.0004 205.916 68.3447 206.115L95.3066 221.682C103.184 226.23 112.889 226.232 120.779 221.679L154.312 202.32C154.271 202.342 154.23 202.362 154.189 202.384C154.283 202.333 154.375 202.281 154.468 202.229L154.312 202.32C154.463 202.237 154.615 202.154 154.765 202.068L154.762 202.065C162.646 197.511 167.487 189.102 167.488 180.009L167.489 89.8618Z" fill="white"/>
<path d="M163.007 36.0259C155.13 31.4781 145.424 31.4763 137.533 36.0288L104.002 55.3882C104.029 55.3732 104.057 55.359 104.084 55.3442C104.02 55.3794 103.956 55.4159 103.893 55.4516L104.002 55.3882C103.851 55.4714 103.699 55.5536 103.549 55.6401L103.552 55.6411C95.6664 60.1948 90.8241 68.6053 90.8232 77.6987L90.8223 167.846C90.8223 169.786 92.8707 171.04 94.5986 170.16L103.523 165.611C116.573 158.959 124.788 145.549 124.788 130.902V62.9897C124.778 58.6479 129.49 55.9211 133.25 58.0698L174.091 81.6479C175.11 82.2363 175.737 83.3238 175.737 84.5005V151.263C175.738 152.716 177.314 153.621 178.568 152.895L196.967 142.28C204.845 137.732 209.704 129.326 209.704 120.221V77.73C209.704 68.6343 204.855 60.2267 196.967 55.6694L190.982 52.2143C190.65 52.0011 190.313 51.792 189.969 51.5932L163.007 36.0259Z" fill="white"/>
</svg>

After

Width:  |  Height:  |  Size: 1.8 KiB

@@ -0,0 +1,4 @@
<svg width="256" height="256" viewBox="0 0 256 256" fill="none" xmlns="http://www.w3.org/2000/svg">
<rect width="256" height="256" transform="matrix(-1 0 0 1 256 0)" fill="white"/>
<path d="M48.6094 179.978V137.487C48.6096 128.382 53.4685 119.976 61.3467 115.428L79.7451 104.813C80.9999 104.087 82.5759 104.992 82.5762 106.445V173.207C82.5762 174.384 83.2036 175.472 84.2227 176.06L125.08 199.647C128.837 201.777 133.533 199.053 133.523 194.718V126.806C133.523 112.388 141.484 99.1684 154.18 92.4135L154.788 92.0971L163.713 87.5473C165.441 86.667 167.489 87.9224 167.489 89.8618L167.488 180.009L167.474 180.86C167.181 189.624 162.399 197.653 154.762 202.065L154.765 202.068C154.615 202.154 154.463 202.237 154.312 202.32L154.468 202.229C154.375 202.281 154.283 202.333 154.189 202.384C154.23 202.362 154.271 202.342 154.312 202.32L120.779 221.679L120.034 222.092C112.534 226.093 103.539 226.092 96.0508 222.095L95.3066 221.682L68.3447 206.115C68.0004 205.916 67.6627 205.707 67.3301 205.494L61.3467 202.039C53.7053 197.624 48.9149 189.595 48.623 180.829L48.6094 179.978ZM175.737 84.5005C175.737 83.3974 175.186 82.3728 174.277 81.7641L174.091 81.6479L133.25 58.0698C129.49 55.9211 124.778 58.6479 124.788 62.9897V130.902L124.782 131.587C124.53 145.966 116.369 159.063 103.523 165.611L94.5986 170.16L94.4355 170.236C92.799 170.938 90.9459 169.803 90.8281 168.026L90.8223 167.846L90.8232 77.6987C90.8241 68.6053 95.6664 60.1948 103.552 55.6411L103.549 55.6401C103.699 55.5536 103.851 55.4714 104.002 55.3882L103.893 55.4516C103.956 55.4159 104.02 55.3794 104.084 55.3442C104.057 55.359 104.029 55.3732 104.002 55.3882L137.533 36.0288C145.424 31.4763 155.13 31.4781 163.007 36.0259L189.969 51.5932C190.313 51.792 190.65 52.0011 190.982 52.2143L196.967 55.6694C204.855 60.2267 209.704 68.6343 209.704 77.73V120.221L209.689 121.073C209.397 129.848 204.599 137.874 196.967 142.28L178.568 152.895C177.314 153.621 175.738 152.716 175.737 151.263V84.5005ZM176.814 149.866L176.819 149.864L176.832 149.856C176.826 149.859 176.82 149.862 176.814 149.866ZM137.023 194.71L137.019 195.038C136.802 201.878 129.341 206.087 123.354 202.692L123.33 202.678L82.4727 179.091C80.3698 177.877 79.0762 175.634 79.0762 173.207V109.239L63.0957 118.458C56.5124 122.259 52.3745 129.183 52.1221 136.752L52.1094 137.487V179.978C52.1094 187.823 56.2909 195.075 63.0957 199.007H63.0967L69.0801 202.462L69.1504 202.503L69.2188 202.547C69.5212 202.741 69.8111 202.92 70.0947 203.083L97.0566 218.651L97.6982 219.007C104.372 222.569 112.434 222.453 119.029 218.648L152.514 199.316L152.512 199.312C152.581 199.274 152.65 199.236 152.752 199.178L152.753 199.181C152.775 199.169 152.796 199.158 152.815 199.147L153.011 199.035C159.81 195.107 163.987 187.854 163.988 180.009V91.3344L156.378 95.2153C144.501 101.27 137.023 113.475 137.023 126.806V194.71ZM94.3223 166.372L101.934 162.493C113.811 156.438 121.288 144.233 121.288 130.902V62.9897C121.278 55.9521 128.901 51.5532 134.986 55.0307L135 55.0385L175.841 78.6167L176.036 78.7339C178.023 79.9714 179.237 82.15 179.237 84.5005V148.468L195.217 139.249C202.013 135.326 206.204 128.075 206.204 120.221V77.73C206.204 69.8848 202.022 62.6318 195.217 58.6997L189.232 55.2456L189.162 55.2046L189.094 55.1606C188.788 54.9649 188.5 54.787 188.219 54.6245L161.257 39.0571C154.463 35.135 146.091 35.1311 139.282 39.0591L139.283 39.06L105.765 58.4106L105.767 58.4135C105.713 58.443 105.723 58.4383 105.602 58.5063L105.6 58.5034C105.535 58.5391 105.481 58.5691 105.433 58.5962L105.302 58.6723C98.5016 62.5995 94.324 69.8532 94.3232 77.6987L94.3223 166.372Z" fill="black"/>
</svg>

After

Width:  |  Height:  |  Size: 3.5 KiB

@@ -0,0 +1,4 @@
<svg width="256" height="256" viewBox="0 0 256 256" fill="none" xmlns="http://www.w3.org/2000/svg">
<rect width="256" height="256" transform="matrix(-1 0 0 1 256 0)" fill="black"/>
<path d="M48.6094 179.978V137.487C48.6096 128.382 53.4685 119.976 61.3467 115.428L79.7451 104.813C80.9999 104.087 82.5759 104.992 82.5762 106.445V173.207C82.5762 174.384 83.2036 175.472 84.2227 176.06L125.08 199.647C128.837 201.777 133.533 199.053 133.523 194.718V126.806C133.523 112.388 141.484 99.1684 154.18 92.4135L154.788 92.0971L163.713 87.5473C165.441 86.667 167.489 87.9224 167.489 89.8618L167.488 180.009L167.474 180.86C167.181 189.624 162.399 197.653 154.762 202.065L154.765 202.068C154.615 202.154 154.463 202.237 154.312 202.32L154.468 202.229C154.375 202.281 154.283 202.333 154.189 202.384C154.23 202.362 154.271 202.342 154.312 202.32L120.779 221.679L120.034 222.092C112.534 226.093 103.539 226.092 96.0508 222.095L95.3066 221.682L68.3447 206.115C68.0004 205.916 67.6627 205.707 67.3301 205.494L61.3467 202.039C53.7053 197.624 48.9149 189.595 48.623 180.829L48.6094 179.978ZM175.737 84.5005C175.737 83.3974 175.186 82.3728 174.277 81.7641L174.091 81.6479L133.25 58.0698C129.49 55.9211 124.778 58.6479 124.788 62.9897V130.902L124.782 131.587C124.53 145.966 116.369 159.063 103.523 165.611L94.5986 170.16L94.4355 170.236C92.799 170.938 90.9459 169.803 90.8281 168.026L90.8223 167.846L90.8232 77.6987C90.8241 68.6053 95.6664 60.1948 103.552 55.6411L103.549 55.6401C103.699 55.5536 103.851 55.4714 104.002 55.3882L103.893 55.4516C103.956 55.4159 104.02 55.3794 104.084 55.3442C104.057 55.359 104.029 55.3732 104.002 55.3882L137.533 36.0288C145.424 31.4763 155.13 31.4781 163.007 36.0259L189.969 51.5932C190.313 51.792 190.65 52.0011 190.982 52.2143L196.967 55.6694C204.855 60.2267 209.704 68.6343 209.704 77.73V120.221L209.689 121.073C209.397 129.848 204.599 137.874 196.967 142.28L178.568 152.895C177.314 153.621 175.738 152.716 175.737 151.263V84.5005ZM176.814 149.866L176.819 149.864L176.832 149.856C176.826 149.859 176.82 149.862 176.814 149.866ZM137.023 194.71L137.019 195.038C136.802 201.878 129.341 206.087 123.354 202.692L123.33 202.678L82.4727 179.091C80.3698 177.877 79.0762 175.634 79.0762 173.207V109.239L63.0957 118.458C56.5124 122.259 52.3745 129.183 52.1221 136.752L52.1094 137.487V179.978C52.1094 187.823 56.2909 195.075 63.0957 199.007H63.0967L69.0801 202.462L69.1504 202.503L69.2188 202.547C69.5212 202.741 69.8111 202.92 70.0947 203.083L97.0566 218.651L97.6982 219.007C104.372 222.569 112.434 222.453 119.029 218.648L152.514 199.316L152.512 199.312C152.581 199.274 152.65 199.236 152.752 199.178L152.753 199.181C152.775 199.169 152.796 199.158 152.815 199.147L153.011 199.035C159.81 195.107 163.987 187.854 163.988 180.009V91.3344L156.378 95.2153C144.501 101.27 137.023 113.475 137.023 126.806V194.71ZM94.3223 166.372L101.934 162.493C113.811 156.438 121.288 144.233 121.288 130.902V62.9897C121.278 55.9521 128.901 51.5532 134.986 55.0307L135 55.0385L175.841 78.6167L176.036 78.7339C178.023 79.9714 179.237 82.15 179.237 84.5005V148.468L195.217 139.249C202.013 135.326 206.204 128.075 206.204 120.221V77.73C206.204 69.8848 202.022 62.6318 195.217 58.6997L189.232 55.2456L189.162 55.2046L189.094 55.1606C188.788 54.9649 188.5 54.787 188.219 54.6245L161.257 39.0571C154.463 35.135 146.091 35.1311 139.282 39.0591L139.283 39.06L105.765 58.4106L105.767 58.4135C105.713 58.443 105.723 58.4383 105.602 58.5063L105.6 58.5034C105.535 58.5391 105.481 58.5691 105.433 58.5962L105.302 58.6723C98.5016 62.5995 94.324 69.8532 94.3232 77.6987L94.3223 166.372Z" fill="white"/>
</svg>

After

Width:  |  Height:  |  Size: 3.5 KiB

@@ -138,7 +138,7 @@ public sealed class HarnessAgent : DelegatingAIAgent
if (options?.DisableToolApproval is not true)
{
builder.UseToolApproval();
builder.UseToolApproval(options?.ToolApprovalAgentOptions);
}
if (options?.DisableOpenTelemetry is not true)
@@ -101,6 +101,15 @@ public sealed class HarnessAgentOptions
/// </remarks>
public bool DisableToolApproval { get; set; }
/// <summary>
/// Gets or sets the options for the <see cref="ToolApprovalAgent"/> middleware.
/// </summary>
/// <remarks>
/// When <see langword="null"/>, the <see cref="ToolApprovalAgent"/> uses default settings.
/// This property has no effect when <see cref="DisableToolApproval"/> is <see langword="true"/>.
/// </remarks>
public ToolApprovalAgentOptions? ToolApprovalAgentOptions { get; set; }
/// <summary>
/// Gets or sets a value indicating whether the <see cref="FileMemoryProvider"/> is disabled.
/// </summary>
@@ -49,7 +49,7 @@ internal sealed class ForeachExecutor : DeclarativeActionExecutor<Foreach>
EvaluationResult<DataValue> expressionResult = this.Evaluator.GetValue(this.Model.Items);
if (expressionResult.Value is TableDataValue tableValue)
{
this._values = [.. tableValue.Values.Select(value => value.Properties.Values.First().ToFormula())];
this._values = [.. tableValue.Values.Select(value => value.ToFormula())];
}
else
{
@@ -181,6 +181,36 @@ public sealed class ChatClientAgentOptions
[Experimental(DiagnosticIds.Experiments.AgentsAIExperiments)]
public bool EnableMessageInjection { get; set; }
/// <summary>
/// Gets or sets a value indicating whether to store automatically approved function calls in the session state
/// for tools that do not require approval when they are returned alongside tools that do.
/// </summary>
/// <remarks>
/// <para>
/// <see cref="FunctionInvokingChatClient"/> has an all-or-nothing behavior for approvals: when any tool
/// in a response is an <see cref="ApprovalRequiredAIFunction"/>, it converts all <see cref="FunctionCallContent"/>
/// items to <see cref="ToolApprovalRequestContent"/>, even for tools that do not require approval.
/// </para>
/// <para>
/// Setting this property to <see langword="true"/> injects an <see cref="NonApprovalRequiredFunctionBypassingChatClient"/>
/// decorator above <see cref="FunctionInvokingChatClient"/> in the pipeline. This decorator identifies approval
/// requests for non-approval-required tools, removes them from the response, and stores them in the session.
/// On the next request, the stored items are automatically re-injected as approved, so the caller only needs
/// to handle approval requests for tools that truly require human approval.
/// </para>
/// <para>
/// This option has no effect when <see cref="UseProvidedChatClientAsIs"/> is <see langword="true"/>.
/// When using a custom chat client stack, you can add an <see cref="NonApprovalRequiredFunctionBypassingChatClient"/>
/// manually via the <see cref="ChatClientBuilderExtensions.UseNonApprovalRequiredFunctionBypassing"/>
/// extension method.
/// </para>
/// </remarks>
/// <value>
/// Default is <see langword="false"/>.
/// </value>
[Experimental(DiagnosticIds.Experiments.AgentsAIExperiments)]
public bool EnableNonApprovalRequiredFunctionBypassing { get; set; }
/// <summary>
/// Creates a new instance of <see cref="ChatClientAgentOptions"/> with the same values as this instance.
/// </summary>
@@ -199,5 +229,6 @@ public sealed class ChatClientAgentOptions
ThrowOnChatHistoryProviderConflict = this.ThrowOnChatHistoryProviderConflict,
RequirePerServiceCallChatHistoryPersistence = this.RequirePerServiceCallChatHistoryPersistence,
EnableMessageInjection = this.EnableMessageInjection,
EnableNonApprovalRequiredFunctionBypassing = this.EnableNonApprovalRequiredFunctionBypassing,
};
}
@@ -148,4 +148,35 @@ public static class ChatClientBuilderExtensions
{
return builder.Use(innerClient => new MessageInjectingChatClient(innerClient));
}
/// <summary>
/// Adds an <see cref="NonApprovalRequiredFunctionBypassingChatClient"/> to the chat client pipeline.
/// </summary>
/// <remarks>
/// <para>
/// This decorator should be positioned above the <see cref="FunctionInvokingChatClient"/> in the pipeline
/// so that it can intercept approval requests for tools that do not require approval. When
/// <see cref="FunctionInvokingChatClient"/> converts all function calls to approval requests (because at
/// least one tool requires approval), this decorator removes the requests for non-approval-required tools,
/// stores them in the session, and automatically re-injects them as approved on the next request.
/// </para>
/// <para>
/// This extension method is intended for use with custom chat client stacks when
/// <see cref="ChatClientAgentOptions.UseProvidedChatClientAsIs"/> is <see langword="true"/>.
/// When <see cref="ChatClientAgentOptions.UseProvidedChatClientAsIs"/> is <see langword="false"/> (the default),
/// the <see cref="ChatClientAgent"/> automatically injects this decorator when
/// <see cref="ChatClientAgentOptions.EnableNonApprovalRequiredFunctionBypassing"/> is <see langword="true"/>.
/// </para>
/// <para>
/// This decorator only works within the context of a running <see cref="ChatClientAgent"/> with
/// an active session, and will throw an exception if used in any other stack.
/// </para>
/// </remarks>
/// <param name="builder">The <see cref="ChatClientBuilder"/> to add the decorator to.</param>
/// <returns>The <paramref name="builder"/> for chaining.</returns>
[Experimental(DiagnosticIds.Experiments.AgentsAIExperiments)]
public static ChatClientBuilder UseNonApprovalRequiredFunctionBypassing(this ChatClientBuilder builder)
{
return builder.Use(innerClient => new NonApprovalRequiredFunctionBypassingChatClient(innerClient));
}
}
@@ -53,6 +53,17 @@ public static class ChatClientExtensions
{
var chatBuilder = chatClient.AsBuilder();
// NonApprovalRequiredFunctionBypassingChatClient is registered before FunctionInvokingChatClient so that
// it sits above FICC in the pipeline. ChatClientBuilder.Build applies factories in reverse order,
// making the first Use() call outermost. By adding this decorator first, the resulting pipeline is:
// NonApprovalRequiredFunctionBypassingChatClient → FunctionInvokingChatClient → ChatHistoryPersistingChatClient → leaf IChatClient
// This allows the decorator to intercept FICC's responses and remove approval requests for tools
// that don't actually require approval, storing them for automatic re-injection on the next request.
if (options?.EnableNonApprovalRequiredFunctionBypassing is true)
{
chatBuilder.Use(innerClient => new NonApprovalRequiredFunctionBypassingChatClient(innerClient));
}
if (chatClient.GetService<FunctionInvokingChatClient>() is null)
{
chatBuilder.Use((innerClient, services) =>
@@ -0,0 +1,285 @@
// Copyright (c) Microsoft. All rights reserved.
using System;
using System.Collections.Generic;
using System.Linq;
using System.Runtime.CompilerServices;
using System.Threading;
using System.Threading.Tasks;
using Microsoft.Extensions.AI;
namespace Microsoft.Agents.AI;
/// <summary>
/// A delegating chat client that automatically removes <see cref="ToolApprovalRequestContent"/> for tools
/// that do not actually require approval, storing auto-approved results in the session for transparent
/// re-injection on the next request.
/// </summary>
/// <remarks>
/// <para>
/// <see cref="FunctionInvokingChatClient"/> has an all-or-nothing behavior for approvals: when any tool
/// in a response is an <see cref="ApprovalRequiredAIFunction"/>, it converts all <see cref="FunctionCallContent"/>
/// items to <see cref="ToolApprovalRequestContent"/> — even for tools that do not require approval. This
/// decorator sits above <see cref="FunctionInvokingChatClient"/> in the pipeline and transparently handles
/// the non-approval-required items so callers only see approval requests for tools that truly need them.
/// </para>
/// <para>
/// On outbound responses, the decorator identifies <see cref="ToolApprovalRequestContent"/> items for tools
/// that are not wrapped in <see cref="ApprovalRequiredAIFunction"/>, removes them from the response, and
/// stores them in the session's <see cref="AgentSessionStateBag"/>. On the next inbound request, the stored
/// items are re-injected as pre-approved <see cref="ToolApprovalResponseContent"/> so that
/// <see cref="FunctionInvokingChatClient"/> can process them alongside the caller's human-approved responses.
/// </para>
/// <para>
/// This decorator requires an active <see cref="AIAgent.CurrentRunContext"/> with a non-null
/// <see cref="AgentRunContext.Session"/>. An <see cref="InvalidOperationException"/> is thrown if no
/// run context or session is available.
/// </para>
/// </remarks>
internal sealed class NonApprovalRequiredFunctionBypassingChatClient : DelegatingChatClient
{
/// <summary>
/// The key used in <see cref="AgentSessionStateBag"/> to store pending auto-approved function calls
/// between agent runs.
/// </summary>
internal const string StateBagKey = "_autoApprovedFunctionCalls";
/// <summary>
/// Initializes a new instance of the <see cref="NonApprovalRequiredFunctionBypassingChatClient"/> class.
/// </summary>
/// <param name="innerClient">The underlying chat client (typically a <see cref="FunctionInvokingChatClient"/>).</param>
public NonApprovalRequiredFunctionBypassingChatClient(IChatClient innerClient)
: base(innerClient)
{
}
/// <inheritdoc/>
public override async Task<ChatResponse> GetResponseAsync(
IEnumerable<ChatMessage> messages,
ChatOptions? options = null,
CancellationToken cancellationToken = default)
{
var session = GetRequiredSession();
var autoApprovableNames = this.GetAutoApprovableToolNames(options);
messages = InjectPendingAutoApprovals(messages, session);
var response = await base.GetResponseAsync(messages, options, cancellationToken).ConfigureAwait(false);
RemoveAutoApprovedFromMessages(response.Messages, autoApprovableNames, session);
return response;
}
/// <inheritdoc/>
public override async IAsyncEnumerable<ChatResponseUpdate> GetStreamingResponseAsync(
IEnumerable<ChatMessage> messages,
ChatOptions? options = null,
[EnumeratorCancellation] CancellationToken cancellationToken = default)
{
var session = GetRequiredSession();
var autoApprovableNames = this.GetAutoApprovableToolNames(options);
messages = InjectPendingAutoApprovals(messages, session);
List<ToolApprovalRequestContent>? autoApproved = null;
try
{
await foreach (var update in base.GetStreamingResponseAsync(messages, options, cancellationToken).ConfigureAwait(false))
{
if (FilterUpdateContents(update, autoApprovableNames, ref autoApproved))
{
yield return update;
}
}
}
finally
{
if (autoApproved is { Count: > 0 })
{
session.StateBag.SetValue(StateBagKey, autoApproved, AgentJsonUtilities.DefaultOptions);
}
}
}
/// <summary>
/// Gets the current <see cref="AgentSession"/> from the ambient run context.
/// </summary>
/// <exception cref="InvalidOperationException">No run context or session is available.</exception>
private static AgentSession GetRequiredSession()
{
var runContext = AIAgent.CurrentRunContext
?? throw new InvalidOperationException(
$"{nameof(NonApprovalRequiredFunctionBypassingChatClient)} can only be used within the context of a running AIAgent. " +
"Ensure that the chat client is being invoked as part of an AIAgent.RunAsync or AIAgent.RunStreamingAsync call.");
return runContext.Session
?? throw new InvalidOperationException(
$"{nameof(NonApprovalRequiredFunctionBypassingChatClient)} requires a session. " +
"Ensure the agent has a resolved session before invoking the chat client.");
}
/// <summary>
/// Checks the session for stored auto-approvals from a previous turn and injects them as
/// a user message containing <see cref="ToolApprovalResponseContent"/> items appended to the input messages.
/// </summary>
/// <remarks>
/// All stored requests are unconditionally injected as approved responses regardless of whether the
/// tool set has changed, because the LLM requires a complete set of tool call responses for a prior turn.
/// </remarks>
private static IEnumerable<ChatMessage> InjectPendingAutoApprovals(
IEnumerable<ChatMessage> messages,
AgentSession session)
{
if (!session.StateBag.TryGetValue<List<ToolApprovalRequestContent>>(
StateBagKey,
out var pendingRequests,
AgentJsonUtilities.DefaultOptions)
|| pendingRequests is not { Count: > 0 })
{
return messages;
}
session.StateBag.TryRemoveValue(StateBagKey);
List<AIContent> approvalResponses = [];
foreach (var request in pendingRequests)
{
approvalResponses.Add(request.CreateResponse(approved: true));
}
var userMessage = new ChatMessage(ChatRole.User, approvalResponses);
return messages.Concat([userMessage]);
}
/// <summary>
/// Builds a set of tool names that do not require approval and can be auto-approved,
/// by checking all available tools from <see cref="ChatOptions.Tools"/> and
/// <see cref="FunctionInvokingChatClient.AdditionalTools"/>.
/// </summary>
private HashSet<string> GetAutoApprovableToolNames(ChatOptions? options)
{
var ficc = this.GetService<FunctionInvokingChatClient>();
var allTools = (options?.Tools ?? Enumerable.Empty<AITool>())
.Concat(ficc?.AdditionalTools ?? Enumerable.Empty<AITool>());
return new HashSet<string>(
allTools
.OfType<AIFunction>()
.Where(static f => f.GetService<ApprovalRequiredAIFunction>() is null)
.Select(static f => f.Name),
StringComparer.Ordinal);
}
/// <summary>
/// Determines whether a <see cref="ToolApprovalRequestContent"/> can be auto-approved because
/// the underlying tool is not an <see cref="ApprovalRequiredAIFunction"/>.
/// </summary>
/// <returns>
/// <see langword="true"/> if the approval request is for a known tool that does not require approval
/// and can be auto-approved; <see langword="false"/> otherwise.
/// </returns>
private static bool IsAutoApprovable(ToolApprovalRequestContent approval, HashSet<string> autoApprovableNames)
{
if (approval.ToolCall is not FunctionCallContent fcc)
{
// Non-function tool calls cannot be auto-approved.
return false;
}
// Auto-approve only if the tool is known and explicitly does NOT require approval.
// Unknown tools are not in the set and are treated as approval-required (safe default).
return autoApprovableNames.Contains(fcc.Name);
}
/// <summary>
/// Scans response messages for auto-approvable <see cref="ToolApprovalRequestContent"/> items,
/// removes them from the messages, and stores them in the session for the next request.
/// </summary>
private static void RemoveAutoApprovedFromMessages(
IList<ChatMessage> messages,
HashSet<string> autoApprovableNames,
AgentSession session)
{
List<ToolApprovalRequestContent>? autoApproved = null;
foreach (var message in messages)
{
for (int i = message.Contents.Count - 1; i >= 0; i--)
{
if (message.Contents[i] is ToolApprovalRequestContent approval
&& IsAutoApprovable(approval, autoApprovableNames))
{
(autoApproved ??= []).Add(approval);
message.Contents.RemoveAt(i);
}
}
}
// Remove messages that are now empty after filtering.
for (int i = messages.Count - 1; i >= 0; i--)
{
if (messages[i].Contents.Count == 0)
{
messages.RemoveAt(i);
}
}
if (autoApproved is { Count: > 0 })
{
session.StateBag.SetValue(StateBagKey, autoApproved, AgentJsonUtilities.DefaultOptions);
}
}
/// <summary>
/// Filters auto-approvable <see cref="ToolApprovalRequestContent"/> items from a streaming update's
/// contents, collecting them for later storage.
/// </summary>
/// <returns>
/// <see langword="true"/> if the update should be yielded (has remaining content or had no
/// approval content to begin with); <see langword="false"/> if the update is now empty and
/// should be skipped.
/// </returns>
private static bool FilterUpdateContents(
ChatResponseUpdate update,
HashSet<string> autoApprovableNames,
ref List<ToolApprovalRequestContent>? autoApproved)
{
bool hasApprovalContent = false;
List<AIContent> filteredContents = [];
bool removedAny = false;
for (int i = 0; i < update.Contents.Count; i++)
{
var content = update.Contents[i];
if (content is ToolApprovalRequestContent approval)
{
hasApprovalContent = true;
if (IsAutoApprovable(approval, autoApprovableNames))
{
(autoApproved ??= []).Add(approval);
removedAny = true;
}
else
{
filteredContents.Add(content);
}
}
else
{
filteredContents.Add(content);
}
}
if (removedAny)
{
update.Contents = filteredContents;
}
// Yield the update unless it was purely auto-approvable approval content (now empty).
return update.Contents.Count > 0 || !hasApprovalContent;
}
}
@@ -51,20 +51,22 @@ public sealed class ToolApprovalAgent : DelegatingAIAgent
{
private readonly ProviderSessionState<ToolApprovalState> _sessionState;
private readonly JsonSerializerOptions _jsonSerializerOptions;
private readonly Func<FunctionCallContent, ValueTask<bool>>[]? _autoApprovalRules;
/// <summary>
/// Initializes a new instance of the <see cref="ToolApprovalAgent"/> class.
/// </summary>
/// <param name="innerAgent">The underlying agent to delegate to.</param>
/// <param name="jsonSerializerOptions">
/// Optional <see cref="JsonSerializerOptions"/> used for serializing argument values when storing rules
/// and for persisting state. When <see langword="null"/>, <see cref="AgentJsonUtilities.DefaultOptions"/> is used.
/// <param name="options">
/// Optional <see cref="ToolApprovalAgentOptions"/> for configuring serialization and auto-approval rules.
/// When <see langword="null"/>, default settings are used.
/// </param>
/// <exception cref="ArgumentNullException"><paramref name="innerAgent"/> is <see langword="null"/>.</exception>
public ToolApprovalAgent(AIAgent innerAgent, JsonSerializerOptions? jsonSerializerOptions = null)
public ToolApprovalAgent(AIAgent innerAgent, ToolApprovalAgentOptions? options = null)
: base(innerAgent)
{
this._jsonSerializerOptions = jsonSerializerOptions ?? AgentJsonUtilities.DefaultOptions;
this._jsonSerializerOptions = options?.JsonSerializerOptions ?? AgentJsonUtilities.DefaultOptions;
this._autoApprovalRules = options?.AutoApprovalRules?.ToArray();
this._sessionState = new ProviderSessionState<ToolApprovalState>(
_ => new ToolApprovalState(),
"toolApprovalState",
@@ -79,7 +81,7 @@ public sealed class ToolApprovalAgent : DelegatingAIAgent
CancellationToken cancellationToken = default)
{
// Steps 1–2: Unwrap AlwaysApprove wrappers, process any queued approval requests.
var (state, callerMessages, nextQueuedItem) = this.PrepareInboundMessages(messages, session);
var (state, callerMessages, nextQueuedItem) = await this.PrepareInboundMessagesAsync(messages, session).ConfigureAwait(false);
if (nextQueuedItem is not null)
{
@@ -98,7 +100,7 @@ public sealed class ToolApprovalAgent : DelegatingAIAgent
var response = await this.InnerAgent.RunAsync(processedMessages, session, options, cancellationToken).ConfigureAwait(false);
// Classify approval requests: auto-approve matching, queue excess, keep first unapproved.
bool allAutoApproved = this.ProcessAndQueueOutboundApprovalRequests(response.Messages, state, session);
bool allAutoApproved = await this.ProcessAndQueueOutboundApprovalRequestsAsync(response.Messages, state, session).ConfigureAwait(false);
if (!allAutoApproved)
{
@@ -119,7 +121,7 @@ public sealed class ToolApprovalAgent : DelegatingAIAgent
[EnumeratorCancellation] CancellationToken cancellationToken = default)
{
// Steps 1–2: Unwrap AlwaysApprove wrappers, process any queued approval requests.
var (state, callerMessages, nextQueuedItem) = this.PrepareInboundMessages(messages, session);
var (state, callerMessages, nextQueuedItem) = await this.PrepareInboundMessagesAsync(messages, session).ConfigureAwait(false);
if (nextQueuedItem is not null)
{
@@ -197,7 +199,7 @@ public sealed class ToolApprovalAgent : DelegatingAIAgent
yield break;
}
// 4. Classify the collected approval requests against standing rules.
// 4. Classify the collected approval requests against standing rules and auto-approval rules.
List<ToolApprovalRequestContent> unapproved = [];
foreach (var tarc in streamedApprovalRequests)
{
@@ -206,6 +208,11 @@ public sealed class ToolApprovalAgent : DelegatingAIAgent
state.CollectedApprovalResponses.Add(
tarc.CreateResponse(approved: true, reason: "Auto-approved by standing rule"));
}
else if (await this.MatchesAutoApprovalRuleAsync(tarc).ConfigureAwait(false))
{
state.CollectedApprovalResponses.Add(
tarc.CreateResponse(approved: true, reason: "Auto-approved by auto-approval rule"));
}
else
{
unapproved.Add(tarc);
@@ -291,9 +298,9 @@ public sealed class ToolApprovalAgent : DelegatingAIAgent
}
/// <summary>
/// Re-evaluates queued approval requests against current rules and auto-approves any that now match.
/// Re-evaluates queued approval requests against current rules and auto-approval rules, and auto-approves any that now match.
/// </summary>
private void DrainAutoApprovableFromQueue(ToolApprovalState state)
private async ValueTask DrainAutoApprovableFromQueueAsync(ToolApprovalState state)
{
for (int i = state.QueuedApprovalRequests.Count - 1; i >= 0; i--)
{
@@ -303,6 +310,12 @@ public sealed class ToolApprovalAgent : DelegatingAIAgent
state.QueuedApprovalRequests[i].CreateResponse(approved: true, reason: "Auto-approved by standing rule"));
state.QueuedApprovalRequests.RemoveAt(i);
}
else if (await this.MatchesAutoApprovalRuleAsync(state.QueuedApprovalRequests[i]).ConfigureAwait(false))
{
state.CollectedApprovalResponses.Add(
state.QueuedApprovalRequests[i].CreateResponse(approved: true, reason: "Auto-approved by auto-approval rule"));
state.QueuedApprovalRequests.RemoveAt(i);
}
}
}
@@ -318,8 +331,8 @@ public sealed class ToolApprovalAgent : DelegatingAIAgent
/// A tuple of (state, processed caller messages, next queued item or <see langword="null"/> if the queue is resolved).
/// When the returned item is non-null, the caller should return/yield it without calling the inner agent.
/// </returns>
private (ToolApprovalState State, List<ChatMessage> CallerMessages, ToolApprovalRequestContent? NextQueuedItem)
PrepareInboundMessages(IEnumerable<ChatMessage> messages, AgentSession? session)
private async ValueTask<(ToolApprovalState State, List<ChatMessage> CallerMessages, ToolApprovalRequestContent? NextQueuedItem)>
PrepareInboundMessagesAsync(IEnumerable<ChatMessage> messages, AgentSession? session)
{
var state = this._sessionState.GetOrInitializeState(session);
@@ -337,7 +350,7 @@ public sealed class ToolApprovalAgent : DelegatingAIAgent
// Re-evaluate remaining queued items — the caller may have added new rules
// (e.g., "always approve this tool") that resolve additional items.
this.DrainAutoApprovableFromQueue(state);
await this.DrainAutoApprovableFromQueueAsync(state).ConfigureAwait(false);
if (state.QueuedApprovalRequests.Count > 0)
{
@@ -386,15 +399,18 @@ public sealed class ToolApprovalAgent : DelegatingAIAgent
/// <see langword="true"/> if all TARc items were auto-approved (caller should re-invoke the inner agent);
/// <see langword="false"/> otherwise.
/// </returns>
private bool ProcessAndQueueOutboundApprovalRequests(
private async ValueTask<bool> ProcessAndQueueOutboundApprovalRequestsAsync(
IList<ChatMessage> responseMessages,
ToolApprovalState state,
AgentSession? session)
{
// Pass 1: Scan all response messages and classify each approval request as
// auto-approved (matches a standing rule) or unapproved (needs caller decision).
var autoApproved = new List<ToolApprovalRequestContent>();
// Pass 1: Scan all response messages and classify each approval request.
// Auto-approved requests (matching a standing rule or auto-approval rule) have their
// responses collected immediately, preserving the original request order, and are
// marked for removal. Unapproved requests are collected for the caller to decide.
var toRemove = new HashSet<ToolApprovalRequestContent>();
var unapproved = new List<ToolApprovalRequestContent>();
int autoApprovedCount = 0;
foreach (var message in responseMessages)
{
@@ -404,7 +420,17 @@ public sealed class ToolApprovalAgent : DelegatingAIAgent
{
if (MatchesRule(tarc, state.Rules, this._jsonSerializerOptions))
{
autoApproved.Add(tarc);
state.CollectedApprovalResponses.Add(
tarc.CreateResponse(approved: true, reason: "Auto-approved by standing rule"));
toRemove.Add(tarc);
autoApprovedCount++;
}
else if (await this.MatchesAutoApprovalRuleAsync(tarc).ConfigureAwait(false))
{
state.CollectedApprovalResponses.Add(
tarc.CreateResponse(approved: true, reason: "Auto-approved by auto-approval rule"));
toRemove.Add(tarc);
autoApprovedCount++;
}
else
{
@@ -415,18 +441,12 @@ public sealed class ToolApprovalAgent : DelegatingAIAgent
}
// Nothing to process: no auto-approved items and at most one unapproved (no queueing needed).
if (autoApproved.Count == 0 && unapproved.Count <= 1)
// No responses were collected above in this case, so state is unmodified and safe to leave.
if (autoApprovedCount == 0 && unapproved.Count <= 1)
{
return false;
}
// Store auto-approved responses for later injection into the inner agent.
foreach (var tarc in autoApproved)
{
state.CollectedApprovalResponses.Add(
tarc.CreateResponse(approved: true, reason: "Auto-approved by standing rule"));
}
// If every approval request was auto-approved, strip them all and signal the caller
// to re-invoke the inner agent immediately with the collected responses.
if (unapproved.Count == 0)
@@ -439,14 +459,10 @@ public sealed class ToolApprovalAgent : DelegatingAIAgent
// Pass 2: Keep only the first unapproved request in the response (for the caller to decide).
// Queue the remaining unapproved requests for subsequent one-at-a-time delivery.
// Remove all auto-approved and queued items from the response messages.
var toRemove = new HashSet<ToolApprovalRequestContent>(autoApproved);
if (unapproved.Count > 1)
for (int i = 1; i < unapproved.Count; i++)
{
for (int i = 1; i < unapproved.Count; i++)
{
toRemove.Add(unapproved[i]);
state.QueuedApprovalRequests.Add(unapproved[i]);
}
toRemove.Add(unapproved[i]);
state.QueuedApprovalRequests.Add(unapproved[i]);
}
// Walk messages in reverse and strip marked items.
@@ -663,8 +679,36 @@ public sealed class ToolApprovalAgent : DelegatingAIAgent
}
/// <summary>
/// Compares stored rule arguments against actual function call arguments for an exact match.
/// Checks whether a <see cref="ToolApprovalRequestContent"/> is approved by any of the configured
/// auto-approval rules (heuristic functions).
/// </summary>
/// <returns>
/// <see langword="true"/> if any auto-approval rule returns <see langword="true"/> for the function call;
/// <see langword="false"/> if no rules are configured, the request is not a function call, or no rule approves it.
/// </returns>
private async ValueTask<bool> MatchesAutoApprovalRuleAsync(ToolApprovalRequestContent request)
{
if (this._autoApprovalRules is not { Length: > 0 })
{
return false;
}
if (request.ToolCall is not FunctionCallContent functionCall)
{
return false;
}
foreach (var rule in this._autoApprovalRules)
{
if (await rule(functionCall).ConfigureAwait(false))
{
return true;
}
}
return false;
}
private static bool ArgumentsMatch(IDictionary<string, string> ruleArguments, IDictionary<string, object?>? callArguments, JsonSerializerOptions jsonSerializerOptions)
{
if (callArguments is null)
@@ -1,7 +1,6 @@
// Copyright (c) Microsoft. All rights reserved.
using System.Diagnostics.CodeAnalysis;
using System.Text.Json;
using Microsoft.Shared.DiagnosticIds;
using Microsoft.Shared.Diagnostics;
@@ -17,9 +16,9 @@ public static class ToolApprovalAgentBuilderExtensions
/// Adds tool approval middleware to the agent pipeline, enabling "don't ask again" approval behavior.
/// </summary>
/// <param name="builder">The <see cref="AIAgentBuilder"/> to which tool approval support will be added.</param>
/// <param name="jsonSerializerOptions">
/// Optional <see cref="JsonSerializerOptions"/> used for serializing argument values when storing rules
/// and for persisting state. When <see langword="null"/>, <see cref="AgentJsonUtilities.DefaultOptions"/> is used.
/// <param name="options">
/// Optional <see cref="ToolApprovalAgentOptions"/> for configuring serialization and auto-approval rules.
/// When <see langword="null"/>, default settings are used.
/// </param>
/// <returns>The <see cref="AIAgentBuilder"/> with tool approval middleware added, enabling method chaining.</returns>
/// <exception cref="System.ArgumentNullException"><paramref name="builder"/> is <see langword="null"/>.</exception>
@@ -32,6 +31,6 @@ public static class ToolApprovalAgentBuilderExtensions
/// </remarks>
public static AIAgentBuilder UseToolApproval(
this AIAgentBuilder builder,
JsonSerializerOptions? jsonSerializerOptions = null)
=> Throw.IfNull(builder).Use(innerAgent => new ToolApprovalAgent(innerAgent, jsonSerializerOptions));
ToolApprovalAgentOptions? options = null)
=> Throw.IfNull(builder).Use(innerAgent => new ToolApprovalAgent(innerAgent, options));
}
@@ -0,0 +1,45 @@
// Copyright (c) Microsoft. All rights reserved.
using System;
using System.Collections.Generic;
using System.Diagnostics.CodeAnalysis;
using System.Text.Json;
using System.Threading.Tasks;
using Microsoft.Extensions.AI;
using Microsoft.Shared.DiagnosticIds;
namespace Microsoft.Agents.AI;
/// <summary>
/// Options for configuring the <see cref="ToolApprovalAgent"/> middleware.
/// </summary>
[Experimental(DiagnosticIds.Experiments.AgentsAIExperiments)]
public class ToolApprovalAgentOptions
{
/// <summary>
/// Gets or sets the <see cref="System.Text.Json.JsonSerializerOptions"/> used for serializing argument values
/// when storing rules and for persisting state.
/// </summary>
/// <remarks>
/// When <see langword="null"/>, <see cref="AgentJsonUtilities.DefaultOptions"/> is used.
/// </remarks>
public JsonSerializerOptions? JsonSerializerOptions { get; set; }
/// <summary>
/// Gets or sets a collection of heuristic functions that can automatically approve function calls
/// that would otherwise require user approval.
/// </summary>
/// <remarks>
/// <para>
/// Each function receives a <see cref="FunctionCallContent"/> representing the tool call that requires approval
/// and returns a <see cref="ValueTask{Boolean}"/> that resolves to <see langword="true"/> to auto-approve
/// the call, or <see langword="false"/> to continue evaluating the next rule.
/// </para>
/// <para>
/// Auto-approval rules are evaluated after standing rules (derived from prior user approvals) but before
/// prompting the user. Rules are evaluated in order; the first rule returning <see langword="true"/>
/// causes the function call to be auto-approved.
/// </para>
/// </remarks>
public IEnumerable<Func<FunctionCallContent, ValueTask<bool>>>? AutoApprovalRules { get; set; }
}
@@ -644,6 +644,51 @@ public class HarnessAgentTests
Assert.Null(agent.GetService<ToolApprovalAgent>());
}
/// <summary>
/// Verify that ToolApprovalAgentOptions auto-approval rules are passed through and actually used.
/// </summary>
[Fact]
public async Task ToolApproval_AutoApprovalRulesAreAppliedAsync()
{
// Arrange — inner client returns an approval request on first call, then final response on second.
var callCount = 0;
var approvalRequest = new ToolApprovalRequestContent("req1", new FunctionCallContent("call1", "ReadTool"));
var mockClient = new Mock<IChatClient>();
mockClient
.Setup(c => c.GetResponseAsync(
It.IsAny<IEnumerable<ChatMessage>>(),
It.IsAny<ChatOptions>(),
It.IsAny<CancellationToken>()))
.ReturnsAsync(() =>
{
callCount++;
if (callCount == 1)
{
return new ChatResponse(new ChatMessage(ChatRole.Assistant, [approvalRequest]));
}
return new ChatResponse(new ChatMessage(ChatRole.Assistant, "Done"));
});
var options = CreateAllDisabledOptions();
options.DisableToolApproval = false;
options.ToolApprovalAgentOptions = new ToolApprovalAgentOptions
{
AutoApprovalRules = [fcc => new ValueTask<bool>(fcc.Name == "ReadTool")]
};
var agent = new HarnessAgent(mockClient.Object, TestMaxContextWindowTokens, TestMaxOutputTokens, options);
var session = await agent.CreateSessionAsync();
// Act
var response = await agent.RunAsync([new ChatMessage(ChatRole.User, "Hi")], session);
// Assert — the auto-approval rule approved the request, so we get "Done" (not an approval request)
Assert.Equal(2, callCount);
Assert.Equal("Done", response.Text);
}
#endregion
#region Feature: OpenTelemetry
@@ -134,6 +134,7 @@ public class ChatClientAgentOptionsTests
ClearOnChatHistoryProviderConflict = false,
WarnOnChatHistoryProviderConflict = false,
ThrowOnChatHistoryProviderConflict = false,
EnableNonApprovalRequiredFunctionBypassing = true,
};
// Act
@@ -150,6 +151,7 @@ public class ChatClientAgentOptionsTests
Assert.Equal(original.ClearOnChatHistoryProviderConflict, clone.ClearOnChatHistoryProviderConflict);
Assert.Equal(original.WarnOnChatHistoryProviderConflict, clone.WarnOnChatHistoryProviderConflict);
Assert.Equal(original.ThrowOnChatHistoryProviderConflict, clone.ThrowOnChatHistoryProviderConflict);
Assert.Equal(original.EnableNonApprovalRequiredFunctionBypassing, clone.EnableNonApprovalRequiredFunctionBypassing);
// ChatOptions should be cloned, not the same reference
Assert.NotSame(original.ChatOptions, clone.ChatOptions);
@@ -0,0 +1,574 @@
// Copyright (c) Microsoft. All rights reserved.
using System;
using System.Collections.Generic;
using System.Linq;
using System.Threading;
using System.Threading.Tasks;
using Microsoft.Extensions.AI;
using Moq;
namespace Microsoft.Agents.AI.UnitTests;
public class NonApprovalRequiredFunctionBypassingChatClientTests
{
#region GetResponseAsync Tests
[Fact]
public async Task GetResponseAsync_NoApprovalContent_PassesThroughUnchangedAsync()
{
// Arrange
var innerClient = CreateMockChatClient((_, _, _) =>
Task.FromResult(new ChatResponse([new ChatMessage(ChatRole.Assistant, "Hello")])));
var decorator = new NonApprovalRequiredFunctionBypassingChatClient(innerClient);
var session = new ChatClientAgentSession();
// Act
var response = await RunWithAgentContextAsync(decorator, session);
// Assert
Assert.Single(response.Messages);
Assert.Equal("Hello", response.Messages[0].Text);
Assert.Equal(0, session.StateBag.Count);
}
[Fact]
public async Task GetResponseAsync_AllToolsRequireApproval_PassesThroughUnchangedAsync()
{
// Arrange
var approvalTool = new ApprovalRequiredAIFunction(AIFunctionFactory.Create(() => "result", "approvalTool"));
var fcc = new FunctionCallContent("call1", "approvalTool");
var approval = new ToolApprovalRequestContent("req1", fcc);
var innerClient = CreateMockChatClient((_, _, _) =>
Task.FromResult(new ChatResponse([new ChatMessage(ChatRole.Assistant, [approval])])));
var decorator = new NonApprovalRequiredFunctionBypassingChatClient(innerClient);
var session = new ChatClientAgentSession();
var options = new ChatOptions { Tools = [approvalTool] };
// Act
var response = await RunWithAgentContextAsync(decorator, session, options);
// Assert — approval request should remain
Assert.Single(response.Messages);
var contents = response.Messages[0].Contents;
Assert.Single(contents);
Assert.IsType<ToolApprovalRequestContent>(contents[0]);
Assert.Equal(0, session.StateBag.Count);
}
[Fact]
public async Task GetResponseAsync_MixedApproval_RemovesNonApprovalItemsAsync()
{
// Arrange
var normalTool = AIFunctionFactory.Create(() => "result", "normalTool");
var approvalTool = new ApprovalRequiredAIFunction(AIFunctionFactory.Create(() => "result", "approvalTool"));
var fccNormal = new FunctionCallContent("call1", "normalTool");
var fccApproval = new FunctionCallContent("call2", "approvalTool");
var approvalNormal = new ToolApprovalRequestContent("req1", fccNormal);
var approvalRequired = new ToolApprovalRequestContent("req2", fccApproval);
var innerClient = CreateMockChatClient((_, _, _) =>
Task.FromResult(new ChatResponse([
new ChatMessage(ChatRole.Assistant, [approvalNormal, approvalRequired])
])));
var decorator = new NonApprovalRequiredFunctionBypassingChatClient(innerClient);
var session = new ChatClientAgentSession();
var options = new ChatOptions { Tools = [normalTool, approvalTool] };
// Act
var response = await RunWithAgentContextAsync(decorator, session, options);
// Assert — only the approval-required item remains in the response
Assert.Single(response.Messages);
var contents = response.Messages[0].Contents;
Assert.Single(contents);
var remainingApproval = Assert.IsType<ToolApprovalRequestContent>(contents[0]);
Assert.Equal("req2", remainingApproval.RequestId);
}
[Fact]
public async Task GetResponseAsync_MixedApproval_StoresAutoApprovedInSessionAsync()
{
// Arrange
var normalTool = AIFunctionFactory.Create(() => "result", "normalTool");
var approvalTool = new ApprovalRequiredAIFunction(AIFunctionFactory.Create(() => "result", "approvalTool"));
var fccNormal = new FunctionCallContent("call1", "normalTool");
var fccApproval = new FunctionCallContent("call2", "approvalTool");
var approvalNormal = new ToolApprovalRequestContent("req1", fccNormal);
var approvalRequired = new ToolApprovalRequestContent("req2", fccApproval);
var innerClient = CreateMockChatClient((_, _, _) =>
Task.FromResult(new ChatResponse([
new ChatMessage(ChatRole.Assistant, [approvalNormal, approvalRequired])
])));
var decorator = new NonApprovalRequiredFunctionBypassingChatClient(innerClient);
var session = new ChatClientAgentSession();
var options = new ChatOptions { Tools = [normalTool, approvalTool] };
// Act
await RunWithAgentContextAsync(decorator, session, options);
// Assert — the auto-approved item should be stored in the session
Assert.True(session.StateBag.TryGetValue<List<ToolApprovalRequestContent>>(
NonApprovalRequiredFunctionBypassingChatClient.StateBagKey, out var stored, AgentJsonUtilities.DefaultOptions));
Assert.NotNull(stored);
Assert.Single(stored!);
Assert.Equal("req1", stored![0].RequestId);
}
[Fact]
public async Task GetResponseAsync_AllNonApproval_RemovesAllApprovalsAndRemovesEmptyMessageAsync()
{
// Arrange
var normalTool = AIFunctionFactory.Create(() => "result", "normalTool");
var fccNormal = new FunctionCallContent("call1", "normalTool");
var approvalNormal = new ToolApprovalRequestContent("req1", fccNormal);
var innerClient = CreateMockChatClient((_, _, _) =>
Task.FromResult(new ChatResponse([
new ChatMessage(ChatRole.Assistant, [approvalNormal])
])));
var decorator = new NonApprovalRequiredFunctionBypassingChatClient(innerClient);
var session = new ChatClientAgentSession();
var options = new ChatOptions { Tools = [normalTool] };
// Act
var response = await RunWithAgentContextAsync(decorator, session, options);
// Assert — the message should be removed since it's now empty
Assert.Empty(response.Messages);
}
[Fact]
public async Task GetResponseAsync_NextRequest_InjectsStoredAutoApprovalsAsync()
{
// Arrange
var fccNormal = new FunctionCallContent("call1", "normalTool");
var storedApproval = new ToolApprovalRequestContent("req1", fccNormal);
var session = new ChatClientAgentSession();
session.StateBag.SetValue(
NonApprovalRequiredFunctionBypassingChatClient.StateBagKey,
new List<ToolApprovalRequestContent> { storedApproval },
AgentJsonUtilities.DefaultOptions);
IEnumerable<ChatMessage>? capturedMessages = null;
var innerClient = CreateMockChatClient((messages, _, _) =>
{
capturedMessages = messages.ToList();
return Task.FromResult(new ChatResponse([new ChatMessage(ChatRole.Assistant, "Done")]));
});
var decorator = new NonApprovalRequiredFunctionBypassingChatClient(innerClient);
var options = new ChatOptions { Tools = [AIFunctionFactory.Create(() => "result", "normalTool")] };
// Act
await RunWithAgentContextAsync(decorator, session, options);
// Assert — the inner client should receive injected messages
Assert.NotNull(capturedMessages);
var messagesList = capturedMessages!.ToList();
// Original user message + user message with approved responses.
Assert.Equal(2, messagesList.Count);
Assert.Equal(ChatRole.User, messagesList[0].Role);
// User message with the auto-approved ToolApprovalResponseContent
Assert.Equal(ChatRole.User, messagesList[1].Role);
var userContent = messagesList[1].Contents.OfType<ToolApprovalResponseContent>().ToList();
Assert.Single(userContent);
Assert.Equal("req1", userContent[0].RequestId);
Assert.True(userContent[0].Approved);
}
[Fact]
public async Task GetResponseAsync_NextRequest_ClearsStoredAfterInjectionAsync()
{
// Arrange
var fccNormal = new FunctionCallContent("call1", "normalTool");
var storedApproval = new ToolApprovalRequestContent("req1", fccNormal);
var session = new ChatClientAgentSession();
session.StateBag.SetValue(
NonApprovalRequiredFunctionBypassingChatClient.StateBagKey,
new List<ToolApprovalRequestContent> { storedApproval },
AgentJsonUtilities.DefaultOptions);
var innerClient = CreateMockChatClient((_, _, _) =>
Task.FromResult(new ChatResponse([new ChatMessage(ChatRole.Assistant, "Done")])));
var decorator = new NonApprovalRequiredFunctionBypassingChatClient(innerClient);
var options = new ChatOptions { Tools = [AIFunctionFactory.Create(() => "result", "normalTool")] };
// Act
await RunWithAgentContextAsync(decorator, session, options);
// Assert — the stored data should be cleared after successful injection
Assert.False(session.StateBag.TryGetValue<List<ToolApprovalRequestContent>>(
NonApprovalRequiredFunctionBypassingChatClient.StateBagKey, out _, AgentJsonUtilities.DefaultOptions));
}
[Fact]
public async Task GetResponseAsync_UnknownTool_TreatedAsApprovalRequiredAsync()
{
// Arrange — tool is not in ChatOptions.Tools
var fccUnknown = new FunctionCallContent("call1", "unknownTool");
var approvalUnknown = new ToolApprovalRequestContent("req1", fccUnknown);
var innerClient = CreateMockChatClient((_, _, _) =>
Task.FromResult(new ChatResponse([
new ChatMessage(ChatRole.Assistant, [approvalUnknown])
])));
var decorator = new NonApprovalRequiredFunctionBypassingChatClient(innerClient);
var session = new ChatClientAgentSession();
var options = new ChatOptions { Tools = [] };
// Act
var response = await RunWithAgentContextAsync(decorator, session, options);
// Assert — unknown tool should NOT be auto-approved
Assert.Single(response.Messages);
Assert.Single(response.Messages[0].Contents);
Assert.IsType<ToolApprovalRequestContent>(response.Messages[0].Contents[0]);
Assert.Equal(0, session.StateBag.Count);
}
[Fact]
public async Task GetResponseAsync_StoredRequestToolSetChanged_StillInjectsAsApprovedAsync()
{
// Arrange — tool was previously non-approval-required but is now wrapped in ApprovalRequiredAIFunction.
// The LLM still requires a complete set of responses, so we inject unconditionally.
var fccTool = new FunctionCallContent("call1", "changingTool");
var storedApproval = new ToolApprovalRequestContent("req1", fccTool);
var session = new ChatClientAgentSession();
session.StateBag.SetValue(
NonApprovalRequiredFunctionBypassingChatClient.StateBagKey,
new List<ToolApprovalRequestContent> { storedApproval },
AgentJsonUtilities.DefaultOptions);
IEnumerable<ChatMessage>? capturedMessages = null;
var innerClient = CreateMockChatClient((messages, _, _) =>
{
capturedMessages = messages.ToList();
return Task.FromResult(new ChatResponse([new ChatMessage(ChatRole.Assistant, "Done")]));
});
var decorator = new NonApprovalRequiredFunctionBypassingChatClient(innerClient);
// The tool is now wrapped in ApprovalRequiredAIFunction — but we still inject unconditionally
var approvalTool = new ApprovalRequiredAIFunction(AIFunctionFactory.Create(() => "result", "changingTool"));
var options = new ChatOptions { Tools = [approvalTool] };
// Act
await RunWithAgentContextAsync(decorator, session, options);
// Assert — the stored request should still be injected as approved
Assert.NotNull(capturedMessages);
var messagesList = capturedMessages!.ToList();
Assert.Equal(2, messagesList.Count);
var userContent = messagesList[1].Contents.OfType<ToolApprovalResponseContent>().ToList();
Assert.Single(userContent);
Assert.Equal("req1", userContent[0].RequestId);
Assert.True(userContent[0].Approved);
// Session should be cleared
Assert.False(session.StateBag.TryGetValue<List<ToolApprovalRequestContent>>(
NonApprovalRequiredFunctionBypassingChatClient.StateBagKey, out _, AgentJsonUtilities.DefaultOptions));
}
#endregion
#region GetStreamingResponseAsync Tests
[Fact]
public async Task GetStreamingResponseAsync_NoApprovalContent_PassesThroughUnchangedAsync()
{
// Arrange
var innerClient = CreateMockStreamingChatClient((_, _, _) =>
ToAsyncEnumerableAsync(
new ChatResponseUpdate(ChatRole.Assistant, "Hello")));
var decorator = new NonApprovalRequiredFunctionBypassingChatClient(innerClient);
var session = new ChatClientAgentSession();
// Act
var updates = new List<ChatResponseUpdate>();
await RunStreamingWithAgentContextAsync(decorator, session, updates);
// Assert
Assert.Single(updates);
Assert.Equal("Hello", updates[0].Text);
Assert.Equal(0, session.StateBag.Count);
}
[Fact]
public async Task GetStreamingResponseAsync_MixedApproval_FiltersNonApprovalItemsAsync()
{
// Arrange
var normalTool = AIFunctionFactory.Create(() => "result", "normalTool");
var approvalTool = new ApprovalRequiredAIFunction(AIFunctionFactory.Create(() => "result", "approvalTool"));
var fccNormal = new FunctionCallContent("call1", "normalTool");
var fccApproval = new FunctionCallContent("call2", "approvalTool");
var approvalNormal = new ToolApprovalRequestContent("req1", fccNormal);
var approvalRequired = new ToolApprovalRequestContent("req2", fccApproval);
var innerClient = CreateMockStreamingChatClient((_, _, _) =>
ToAsyncEnumerableAsync(
new ChatResponseUpdate(ChatRole.Assistant, "text"),
new ChatResponseUpdate { Contents = [approvalNormal, approvalRequired] }));
var decorator = new NonApprovalRequiredFunctionBypassingChatClient(innerClient);
var session = new ChatClientAgentSession();
var options = new ChatOptions { Tools = [normalTool, approvalTool] };
// Act
var updates = new List<ChatResponseUpdate>();
await RunStreamingWithAgentContextAsync(decorator, session, updates, options);
// Assert — text update + filtered approval update
Assert.Equal(2, updates.Count);
Assert.Equal("text", updates[0].Text);
// Second update should only have the approval-required item
var approvalContents = updates[1].Contents.OfType<ToolApprovalRequestContent>().ToList();
Assert.Single(approvalContents);
Assert.Equal("req2", approvalContents[0].RequestId);
}
[Fact]
public async Task GetStreamingResponseAsync_MixedApproval_StoresAutoApprovedInSessionAsync()
{
// Arrange
var normalTool = AIFunctionFactory.Create(() => "result", "normalTool");
var approvalTool = new ApprovalRequiredAIFunction(AIFunctionFactory.Create(() => "result", "approvalTool"));
var fccNormal = new FunctionCallContent("call1", "normalTool");
var fccApproval = new FunctionCallContent("call2", "approvalTool");
var approvalNormal = new ToolApprovalRequestContent("req1", fccNormal);
var approvalRequired = new ToolApprovalRequestContent("req2", fccApproval);
var innerClient = CreateMockStreamingChatClient((_, _, _) =>
ToAsyncEnumerableAsync(
new ChatResponseUpdate { Contents = [approvalNormal, approvalRequired] }));
var decorator = new NonApprovalRequiredFunctionBypassingChatClient(innerClient);
var session = new ChatClientAgentSession();
var options = new ChatOptions { Tools = [normalTool, approvalTool] };
// Act
var updates = new List<ChatResponseUpdate>();
await RunStreamingWithAgentContextAsync(decorator, session, updates, options);
// Assert — the auto-approved item should be stored in the session
Assert.True(session.StateBag.TryGetValue<List<ToolApprovalRequestContent>>(
NonApprovalRequiredFunctionBypassingChatClient.StateBagKey, out var stored, AgentJsonUtilities.DefaultOptions));
Assert.NotNull(stored);
Assert.Single(stored!);
Assert.Equal("req1", stored![0].RequestId);
}
[Fact]
public async Task GetStreamingResponseAsync_AllNonApproval_SkipsEmptyUpdateAsync()
{
// Arrange
var normalTool = AIFunctionFactory.Create(() => "result", "normalTool");
var fccNormal = new FunctionCallContent("call1", "normalTool");
var approvalNormal = new ToolApprovalRequestContent("req1", fccNormal);
var innerClient = CreateMockStreamingChatClient((_, _, _) =>
ToAsyncEnumerableAsync(
new ChatResponseUpdate(ChatRole.Assistant, "text"),
new ChatResponseUpdate { Contents = [approvalNormal] }));
var decorator = new NonApprovalRequiredFunctionBypassingChatClient(innerClient);
var session = new ChatClientAgentSession();
var options = new ChatOptions { Tools = [normalTool] };
// Act
var updates = new List<ChatResponseUpdate>();
await RunStreamingWithAgentContextAsync(decorator, session, updates, options);
// Assert — the approval update should be skipped entirely
Assert.Single(updates);
Assert.Equal("text", updates[0].Text);
}
#endregion
#region Error Handling Tests
[Fact]
public async Task GetResponseAsync_NoRunContext_ThrowsInvalidOperationExceptionAsync()
{
// Arrange
var innerClient = CreateMockChatClient((_, _, _) =>
Task.FromResult(new ChatResponse([new ChatMessage(ChatRole.Assistant, "response")])));
var decorator = new NonApprovalRequiredFunctionBypassingChatClient(innerClient);
// Act & Assert — calling directly without agent context
await Assert.ThrowsAsync<InvalidOperationException>(
() => decorator.GetResponseAsync([new ChatMessage(ChatRole.User, "test")]));
}
[Fact]
public async Task GetResponseAsync_NoSession_ThrowsInvalidOperationExceptionAsync()
{
// Arrange
var innerClient = CreateMockChatClient((_, _, _) =>
Task.FromResult(new ChatResponse([new ChatMessage(ChatRole.Assistant, "response")])));
var decorator = new NonApprovalRequiredFunctionBypassingChatClient(innerClient);
// Act & Assert — run with null session
await Assert.ThrowsAsync<InvalidOperationException>(
() => RunWithAgentContextAsync(decorator, session: null!));
}
#endregion
#region Builder Extension Tests
[Fact]
public void UseNonApprovalRequiredFunctionBypassing_AddsDecoratorToPipeline()
{
// Arrange
var innerClient = new Mock<IChatClient>().Object;
// Act
var pipeline = innerClient.AsBuilder()
.UseNonApprovalRequiredFunctionBypassing()
.Build();
// Assert
Assert.NotNull(pipeline.GetService<NonApprovalRequiredFunctionBypassingChatClient>());
}
[Fact]
public void WithDefaultAgentMiddleware_EnableNonApprovalRequiredFunctionBypassing_InjectsDecorator()
{
// Arrange
var innerClient = new Mock<IChatClient>().Object;
var options = new ChatClientAgentOptions { EnableNonApprovalRequiredFunctionBypassing = true };
// Act
var pipeline = innerClient.WithDefaultAgentMiddleware(options);
// Assert
Assert.NotNull(pipeline.GetService<NonApprovalRequiredFunctionBypassingChatClient>());
}
[Fact]
public void WithDefaultAgentMiddleware_EnableNonApprovalRequiredFunctionBypassingFalse_DoesNotInjectDecorator()
{
// Arrange
var innerClient = new Mock<IChatClient>().Object;
var options = new ChatClientAgentOptions { EnableNonApprovalRequiredFunctionBypassing = false };
// Act
var pipeline = innerClient.WithDefaultAgentMiddleware(options);
// Assert
Assert.Null(pipeline.GetService<NonApprovalRequiredFunctionBypassingChatClient>());
}
#endregion
#region Helpers
private static async Task<ChatResponse> RunWithAgentContextAsync(
NonApprovalRequiredFunctionBypassingChatClient decorator,
AgentSession? session,
ChatOptions? options = null)
{
ChatResponse? capturedResponse = null;
var agent = new TestAIAgent
{
RunAsyncFunc = async (messages, agentSession, agentOptions, ct) =>
{
capturedResponse = await decorator.GetResponseAsync(messages, options, ct);
return new AgentResponse(capturedResponse);
}
};
await agent.RunAsync([new ChatMessage(ChatRole.User, "Hello")], session);
return capturedResponse!;
}
private static Task<ChatResponse> RunWithAgentContextAsync(
NonApprovalRequiredFunctionBypassingChatClient decorator,
AgentSession session)
=> RunWithAgentContextAsync(decorator, session, options: null);
private static async Task RunStreamingWithAgentContextAsync(
NonApprovalRequiredFunctionBypassingChatClient decorator,
AgentSession session,
List<ChatResponseUpdate> updates,
ChatOptions? options = null)
{
var agent = new TestAIAgent
{
RunAsyncFunc = async (messages, agentSession, agentOptions, ct) =>
{
await foreach (var update in decorator.GetStreamingResponseAsync(messages, options, ct))
{
updates.Add(update);
}
return new AgentResponse([new ChatMessage(ChatRole.Assistant, "done")]);
}
};
await agent.RunAsync([new ChatMessage(ChatRole.User, "Hello")], session);
}
private static IChatClient CreateMockChatClient(
Func<IEnumerable<ChatMessage>, ChatOptions?, CancellationToken, Task<ChatResponse>> onGetResponse)
{
var mock = new Mock<IChatClient>();
mock.Setup(c => c.GetResponseAsync(
It.IsAny<IEnumerable<ChatMessage>>(),
It.IsAny<ChatOptions?>(),
It.IsAny<CancellationToken>()))
.Returns((IEnumerable<ChatMessage> m, ChatOptions? o, CancellationToken ct) => onGetResponse(m, o, ct));
return mock.Object;
}
private static IChatClient CreateMockStreamingChatClient(
Func<IEnumerable<ChatMessage>, ChatOptions?, CancellationToken, IAsyncEnumerable<ChatResponseUpdate>> onGetStreamingResponse)
{
var mock = new Mock<IChatClient>();
mock.Setup(c => c.GetStreamingResponseAsync(
It.IsAny<IEnumerable<ChatMessage>>(),
It.IsAny<ChatOptions?>(),
It.IsAny<CancellationToken>()))
.Returns((IEnumerable<ChatMessage> m, ChatOptions? o, CancellationToken ct) => onGetStreamingResponse(m, o, ct));
return mock.Object;
}
private static async IAsyncEnumerable<ChatResponseUpdate> ToAsyncEnumerableAsync(params ChatResponseUpdate[] updates)
{
foreach (var update in updates)
{
yield return update;
}
await Task.CompletedTask;
}
#endregion
}
@@ -59,15 +59,15 @@ public class ToolApprovalAgentBuilderExtensionsTests
/// Verify that UseToolApproval with custom JsonSerializerOptions works correctly.
/// </summary>
[Fact]
public void UseToolApproval_WithCustomJsonSerializerOptions_ReturnsToolApprovalAgent()
public void UseToolApproval_WithCustomOptions_ReturnsToolApprovalAgent()
{
// Arrange
var mockAgent = new Mock<AIAgent>();
var builder = new AIAgentBuilder(mockAgent.Object);
var options = new JsonSerializerOptions();
var options = new ToolApprovalAgentOptions { JsonSerializerOptions = new JsonSerializerOptions() };
// Act
var result = builder.UseToolApproval(jsonSerializerOptions: options).Build();
var result = builder.UseToolApproval(options: options).Build();
// Assert
Assert.IsType<ToolApprovalAgent>(result);
@@ -47,14 +47,14 @@ public class ToolApprovalAgentTests
}
/// <summary>
/// Verify that constructor accepts custom JsonSerializerOptions.
/// Verify that constructor accepts custom options.
/// </summary>
[Fact]
public void Constructor_CustomJsonSerializerOptions_CreatesInstanceAsync()
public void Constructor_CustomOptions_CreatesInstance()
{
// Arrange
var innerAgent = new Mock<AIAgent>().Object;
var options = new JsonSerializerOptions();
var options = new ToolApprovalAgentOptions { JsonSerializerOptions = new JsonSerializerOptions() };
// Act
var agent = new ToolApprovalAgent(innerAgent, options);
@@ -1535,4 +1535,311 @@ public class ToolApprovalAgentTests
}
#endregion
#region Auto-Approval Rules (Heuristics)
/// <summary>
/// Verify that an auto-approval rule can approve a function call that would otherwise need user approval.
/// </summary>
[Fact]
public async Task RunAsync_AutoApprovalRule_ApprovesMatchingToolAsync()
{
// Arrange
var session = new ChatClientAgentSession();
var approvalRequest = new ToolApprovalRequestContent("req1", new FunctionCallContent("call1", "ReadTool"));
// Inner agent: first call returns approval request, second returns final response.
var callCount = 0;
var innerAgent = new Mock<AIAgent>();
innerAgent
.Protected()
.Setup<Task<AgentResponse>>("RunCoreAsync",
ItExpr.IsAny<IEnumerable<ChatMessage>>(),
ItExpr.IsAny<AgentSession?>(),
ItExpr.IsAny<AgentRunOptions?>(),
ItExpr.IsAny<CancellationToken>())
.ReturnsAsync(() =>
{
callCount++;
if (callCount == 1)
{
return new AgentResponse([new ChatMessage(ChatRole.Assistant, [approvalRequest])]);
}
return new AgentResponse([new ChatMessage(ChatRole.Assistant, "Done")]);
});
var options = new ToolApprovalAgentOptions
{
AutoApprovalRules = [fcc => new ValueTask<bool>(fcc.Name == "ReadTool")]
};
var agent = new ToolApprovalAgent(innerAgent.Object, options);
// Act
var response = await agent.RunAsync(
[new ChatMessage(ChatRole.User, "Hi")],
session);
// Assert — the approval request was auto-approved, inner agent called twice
Assert.Equal(2, callCount);
Assert.Equal("Done", response.Text);
}
/// <summary>
/// Verify that when auto-approval rule does not match, request is surfaced to the caller.
/// </summary>
[Fact]
public async Task RunAsync_AutoApprovalRule_DoesNotMatchSurfacesToCallerAsync()
{
// Arrange
var session = new ChatClientAgentSession();
var approvalRequest = new ToolApprovalRequestContent("req1", new FunctionCallContent("call1", "DangerousTool"));
var innerAgent = CreateMockAgent(new AgentResponse([new ChatMessage(ChatRole.Assistant, [approvalRequest])]));
var options = new ToolApprovalAgentOptions
{
AutoApprovalRules = [fcc => new ValueTask<bool>(fcc.Name == "ReadTool")] // Only approves ReadTool
};
var agent = new ToolApprovalAgent(innerAgent.Object, options);
// Act
var response = await agent.RunAsync(
[new ChatMessage(ChatRole.User, "Hi")],
session);
// Assert — request surfaced to caller since heuristic doesn't match
var requests = response.Messages.SelectMany(m => m.Contents).OfType<ToolApprovalRequestContent>().ToList();
Assert.Single(requests);
Assert.Equal("DangerousTool", ((FunctionCallContent)requests[0].ToolCall).Name);
}
/// <summary>
/// Verify that multiple auto-approval rules are evaluated in order; first match wins.
/// </summary>
[Fact]
public async Task RunAsync_MultipleAutoApprovalRules_FirstMatchWinsAsync()
{
// Arrange
var session = new ChatClientAgentSession();
var approvalRequest = new ToolApprovalRequestContent("req1", new FunctionCallContent("call1", "SpecialTool"));
var callCount = 0;
var innerAgent = new Mock<AIAgent>();
innerAgent
.Protected()
.Setup<Task<AgentResponse>>("RunCoreAsync",
ItExpr.IsAny<IEnumerable<ChatMessage>>(),
ItExpr.IsAny<AgentSession?>(),
ItExpr.IsAny<AgentRunOptions?>(),
ItExpr.IsAny<CancellationToken>())
.ReturnsAsync(() =>
{
callCount++;
if (callCount == 1)
{
return new AgentResponse([new ChatMessage(ChatRole.Assistant, [approvalRequest])]);
}
return new AgentResponse([new ChatMessage(ChatRole.Assistant, "Done")]);
});
var rule1Called = false;
var rule2Called = false;
var options = new ToolApprovalAgentOptions
{
AutoApprovalRules =
[
fcc => { rule1Called = true; return new ValueTask<bool>(fcc.Name == "SpecialTool"); },
fcc => { rule2Called = true; return new ValueTask<bool>(true); } // Should not be reached
]
};
var agent = new ToolApprovalAgent(innerAgent.Object, options);
// Act
await agent.RunAsync([new ChatMessage(ChatRole.User, "Hi")], session);
// Assert — first rule matched, second was never called
Assert.True(rule1Called);
Assert.False(rule2Called);
}
/// <summary>
/// Verify that standing rules are evaluated before auto-approval rules.
/// </summary>
[Fact]
public async Task RunAsync_StandingRuleTakesPrecedenceOverAutoApprovalRuleAsync()
{
// Arrange
var session = new ChatClientAgentSession();
var approvalRequest = new ToolApprovalRequestContent("req1", new FunctionCallContent("call1", "MyTool"));
var callCount = 0;
var innerAgent = new Mock<AIAgent>();
innerAgent
.Protected()
.Setup<Task<AgentResponse>>("RunCoreAsync",
ItExpr.IsAny<IEnumerable<ChatMessage>>(),
ItExpr.IsAny<AgentSession?>(),
ItExpr.IsAny<AgentRunOptions?>(),
ItExpr.IsAny<CancellationToken>())
.ReturnsAsync(() =>
{
callCount++;
if (callCount <= 2)
{
return new AgentResponse([new ChatMessage(ChatRole.Assistant, [approvalRequest])]);
}
return new AgentResponse([new ChatMessage(ChatRole.Assistant, "Done")]);
});
var heuristicCalled = false;
var options = new ToolApprovalAgentOptions
{
AutoApprovalRules = [fcc => { heuristicCalled = true; return new ValueTask<bool>(true); }]
};
var agent = new ToolApprovalAgent(innerAgent.Object, options);
// Call 1: heuristic should be called (no standing rule yet)
var response1 = await agent.RunAsync([new ChatMessage(ChatRole.User, "Hi")], session);
Assert.True(heuristicCalled);
Assert.Equal("Done", response1.Text);
// Now establish a standing rule by sending AlwaysApprove
heuristicCalled = false;
callCount = 0;
var alwaysApprove = new AlwaysApproveToolApprovalResponseContent(
approvalRequest.CreateResponse(approved: true),
alwaysApproveTool: true,
alwaysApproveToolWithArguments: false);
// Call 2: standing rule should match first, heuristic should NOT be called
var response2 = await agent.RunAsync(
[new ChatMessage(ChatRole.User, [alwaysApprove])],
session);
Assert.False(heuristicCalled);
Assert.Equal("Done", response2.Text);
}
/// <summary>
/// Verify that when a batch contains a mix of heuristic-approved and standing-rule-approved
/// requests, the collected approval responses preserve the original request order rather than
/// being grouped by approval kind.
/// </summary>
[Fact]
public async Task RunAsync_MixedAutoApprovals_PreserveOriginalOrderAsync()
{
// Arrange
var session = new ChatClientAgentSession();
// Batch ordering: first request is approved by a heuristic, second by a standing rule.
var heuristicRequest = new ToolApprovalRequestContent("reqA", new FunctionCallContent("callA", "HeuristicTool"));
var standingRequest = new ToolApprovalRequestContent("reqB", new FunctionCallContent("callB", "StandingTool"));
var batchResponse = new AgentResponse([new ChatMessage(ChatRole.Assistant, [heuristicRequest, standingRequest])]);
var finalResponse = new AgentResponse([new ChatMessage(ChatRole.Assistant, "Done")]);
var callCount = 0;
List<ChatMessage>? secondCallMessages = null;
var innerAgent = new Mock<AIAgent>();
innerAgent
.Protected()
.Setup<Task<AgentResponse>>("RunCoreAsync",
ItExpr.IsAny<IEnumerable<ChatMessage>>(),
ItExpr.IsAny<AgentSession?>(),
ItExpr.IsAny<AgentRunOptions?>(),
ItExpr.IsAny<CancellationToken>())
.Callback<IEnumerable<ChatMessage>, AgentSession?, AgentRunOptions?, CancellationToken>((msgs, _, _, _) =>
{
callCount++;
if (callCount == 2)
{
secondCallMessages = msgs.ToList();
}
})
.ReturnsAsync(() => callCount == 1 ? batchResponse : finalResponse);
var options = new ToolApprovalAgentOptions
{
AutoApprovalRules = [fcc => new ValueTask<bool>(fcc.Name == "HeuristicTool")]
};
var agent = new ToolApprovalAgent(innerAgent.Object, options);
// Establish a standing rule for "StandingTool" via an AlwaysApprove response in the same call.
var alwaysApprove = standingRequest.CreateAlwaysApproveToolResponse("User said always");
// Act — both requests auto-approve (heuristic + standing rule), so the inner agent is re-invoked.
var response = await agent.RunAsync(
[new ChatMessage(ChatRole.User, [alwaysApprove])],
session);
// Assert — inner agent re-called and final response returned.
Assert.Equal(2, callCount);
Assert.Equal("Done", response.Text);
// The injected approval responses must preserve the original request order: reqA before reqB,
// even though reqA was approved by a heuristic and reqB by a standing rule.
Assert.NotNull(secondCallMessages);
var injected = secondCallMessages!
.SelectMany(m => m.Contents)
.OfType<ToolApprovalResponseContent>()
.Where(r => r.RequestId is "reqA" or "reqB")
.ToList();
Assert.Equal(2, injected.Count);
Assert.Equal("reqA", injected[0].RequestId);
Assert.Equal("reqB", injected[1].RequestId);
Assert.All(injected, r => Assert.True(r.Approved));
}
/// <summary>
/// Verify that auto-approval rules work in the streaming path.
/// </summary>
[Fact]
public async Task RunStreamingAsync_AutoApprovalRule_ApprovesMatchingToolAsync()
{
// Arrange
var session = new ChatClientAgentSession();
var approvalRequest = new ToolApprovalRequestContent("req1", new FunctionCallContent("call1", "ReadTool"));
var callCount = 0;
var innerAgent = new Mock<AIAgent>();
innerAgent
.Protected()
.Setup<IAsyncEnumerable<AgentResponseUpdate>>("RunCoreStreamingAsync",
ItExpr.IsAny<IEnumerable<ChatMessage>>(),
ItExpr.IsAny<AgentSession?>(),
ItExpr.IsAny<AgentRunOptions?>(),
ItExpr.IsAny<CancellationToken>())
.Returns(() =>
{
callCount++;
if (callCount == 1)
{
return ToAsyncEnumerableAsync([new AgentResponseUpdate(ChatRole.Assistant, [approvalRequest])]);
}
return ToAsyncEnumerableAsync([new AgentResponseUpdate(ChatRole.Assistant, "Done")]);
});
var options = new ToolApprovalAgentOptions
{
AutoApprovalRules = [fcc => new ValueTask<bool>(fcc.Name == "ReadTool")]
};
var agent = new ToolApprovalAgent(innerAgent.Object, options);
// Act
var updates = new List<AgentResponseUpdate>();
await foreach (var update in agent.RunStreamingAsync([new ChatMessage(ChatRole.User, "Hi")], session))
{
updates.Add(update);
}
// Assert — the approval request was auto-approved, inner agent streamed twice
Assert.Equal(2, callCount);
Assert.Single(updates);
Assert.Equal("Done", updates[0].Text);
}
#endregion
}
@@ -142,6 +142,34 @@ public sealed class ForeachExecutorTest(ITestOutputHelper output) : WorkflowActi
indexName: "CurrentIndex");
}
[Fact]
public async Task ForeachTakeNextWithMultiFieldRecordAsync()
{
// Arrange
const string CurrentValueName = "CurrentValue";
this.SetVariableState(CurrentValueName);
TableDataValue tableValue = DataValue.TableFromRecords(
DataValue.RecordFromFields(
new KeyValuePair<string, DataValue>("name", new StringDataValue("Alice")),
new KeyValuePair<string, DataValue>("role", new StringDataValue("Engineer"))));
Foreach model = this.CreateModel(
displayName: nameof(ForeachTakeNextWithMultiFieldRecordAsync),
items: ValueExpression.Literal(tableValue),
valueName: CurrentValueName,
indexName: null);
ForeachExecutor action = new(model, this.State);
// Act
await this.ExecuteAsync(action, ForeachExecutor.Steps.Next(action.Id), action.TakeNextAsync);
// Assert
RecordValue currentValue = Assert.IsType<RecordValue>(this.State.Get(CurrentValueName), exactMatch: false);
Assert.Equal("Alice", currentValue.GetField("name").ToObject());
Assert.Equal("Engineer", currentValue.GetField("role").ToObject());
}
[Fact]
public async Task ForeachTakeLastAsync()
{
+13
View File
@@ -76,6 +76,19 @@ agent_framework/
- **`SkillScriptRunner`** - Protocol for file-based script execution. Any callable matching `(skill, script, args) -> Any` satisfies it. Code-defined scripts do not use a runner.
- **`SkillsProvider`** - Context provider (extends `ContextProvider`) that discovers file-based skills from `SKILL.md` files and/or accepts code-defined `Skill` instances. Follows progressive disclosure: advertise → load → read resources / run scripts.
### Model Context Protocol (`_mcp.py`)
- **`MCPTool`** - Base wrapper that owns the MCP `ClientSession` and exposes the remote server's tools as `FunctionTool`s.
- **`MCPStdioTool`** / **`MCPStreamableHTTPTool`** / **`MCPWebsocketTool`** - Transport-specific subclasses.
- **`MCPTaskOptions`** (experimental, `MCP_LONG_RUNNING_TASKS` feature, **frozen**) - Per-tool-instance options controlling the SEP-2663 long-running task lifecycle. When the server advertises a tool with `execution.taskSupport == "required"`, `MCPTool.call_tool` transparently routes through `call_tool_as_task`, which sends an augmented `tools/call`, polls `tasks/get` until terminal, and reinterprets `tasks/result` as a normal `CallToolResult`. Instances are immutable; replace via `MCPTool.task_options = MCPTaskOptions(...)`. Fields:
- `default_ttl: timedelta | None` — forwarded to the server as `params.task.ttl` (milliseconds). When `None`, the server's default applies.
- `cancel_remote_task_on_local_cancellation: bool = True` — only gates the `CancelledError` path. Abandonment paths (see below) always cancel.
- `max_task_wait: timedelta | None` — client-side deadline for the whole post-create lifecycle (poll + result fetch). When exceeded, raises `ToolExecutionException` and fires a best-effort `tasks/cancel`. `None` (default) means no client-side bound. Bounds sleeps, sends, AND reconnects via `asyncio.wait_for`.
- **Permissive fallback**: servers that ignore the augmentation (return `CallToolResult` directly) or reject the unknown `task` field with `METHOD_NOT_FOUND` / `INVALID_PARAMS` fall back to the plain `session.call_tool(...)` path so legacy servers keep working. An unparseable success response (server accepted the augmented call but returned a payload that is neither `CreateTaskResult` nor `CallToolResult`) **does not** fall back — it raises `ToolExecutionException` to avoid double-executing a side-effecting tool.
- **Submit-vs-track reconnect policy**: a dropped connection before a `task_id` is known raises `ToolExecutionException("connection lost; task state unknown")` without re-issuing the augmented `tools/call`, so a server that accepted the request but lost the response cannot be made to start the same operation twice; once a `task_id` exists, `tasks/get` / `tasks/result` reconnect once and retry against the same id (a shared `_send_with_one_reconnect` helper).
- **Cancel-on-abandonment vs terminal failure**: any path where the remote task may still be running (max-wait exceeded, hard `McpError` in poll, malformed `tasks/get`, second connection loss in poll/fetch, reconnect failure) fires best-effort `tasks/cancel` before raising. Terminal failures (`failed`/`cancelled`/`input_required` server-side, `completed+isError`, malformed `tasks/result` after server completed) do **not** cancel — the server is already done. `_MCPTaskAbandoned` is the private marker distinguishing the two.
- **Transient poll retry**: a slow `tasks/get` that surfaces as `McpError(code=408 REQUEST_TIMEOUT)` is retried (bounded by `max_task_wait`). All other non-connection `McpError`s during poll are treated as abandonment. `tasks/result` does not get transient retry — the server has already completed, so a slow payload fetch is anomalous.
### File Access Harness (`_harness/_file_access.py`)
- **`AgentFileStore`** - Abstract async store backing the file-access harness. Implementations expose `write_file`, `read_file`, `delete_file`, `list_files`, `file_exists`, `search_files`, and `create_directory` over forward-slash relative paths.
@@ -124,7 +124,7 @@ from ._harness._todo import (
TodoSessionStore,
TodoStore,
)
from ._mcp import MCPStdioTool, MCPStreamableHTTPTool, MCPWebsocketTool
from ._mcp import MCPStdioTool, MCPStreamableHTTPTool, MCPTaskOptions, MCPWebsocketTool
from ._middleware import (
AgentContext,
AgentMiddleware,
@@ -444,12 +444,13 @@ __all__ = [
"InlineSkillResource",
"InlineSkillScript",
"LocalEvaluator",
"MCPStdioTool",
"MCPStreamableHTTPTool",
"MCPWebsocketTool",
"MCPSkill",
"MCPSkillResource",
"MCPSkillsSource",
"MCPStdioTool",
"MCPStreamableHTTPTool",
"MCPTaskOptions",
"MCPWebsocketTool",
"MemoryContextProvider",
"MemoryFileStore",
"MemoryIndexEntry",
@@ -58,6 +58,7 @@ class ExperimentalFeature(str, Enum):
FOUNDRY_PREVIEW_TOOLS = "FOUNDRY_PREVIEW_TOOLS"
FUNCTIONAL_WORKFLOWS = "FUNCTIONAL_WORKFLOWS"
HARNESS = "HARNESS"
MCP_LONG_RUNNING_TASKS = "MCP_LONG_RUNNING_TASKS"
MCP_SKILLS = "MCP_SKILLS"
PROGRESSIVE_TOOLS = "PROGRESSIVE_TOOLS"
SKILLS = "SKILLS"
File diff suppressed because it is too large Load Diff
@@ -2,7 +2,6 @@
from __future__ import annotations
import json
import logging
import sys
import uuid
@@ -12,7 +11,6 @@ from datetime import datetime, timezone
from typing import TYPE_CHECKING, Any, ClassVar, Literal, cast, overload
from .._agents import BaseAgent
from .._serialization import make_json_safe
from .._sessions import (
AgentSession,
ContextProvider,
@@ -30,11 +28,12 @@ from .._types import (
UsageDetails,
add_usage_details,
)
from ..exceptions import AgentInvalidRequestException, AgentInvalidResponseException
from ..exceptions import AgentException, AgentInvalidRequestException, AgentInvalidResponseException
from ._checkpoint import CheckpointStorage
from ._events import (
AGENT_FORWARDED_EVENT_TYPES,
WorkflowEvent,
WorkflowRunState,
)
from ._message_utils import normalize_messages_input
from ._typing_utils import is_instance_of, is_type_compatible
@@ -59,27 +58,24 @@ class WorkflowAgent(BaseAgent):
@dataclass
class RequestInfoFunctionArgs:
request_id: str
data: Any
request_event: WorkflowEvent
def to_dict(self) -> dict[str, Any]:
return {"request_id": self.request_id, "data": make_json_safe(self.data)}
def to_json(self) -> str:
return json.dumps(self.to_dict())
return {"request_id": self.request_id, "request_event": self.request_event.to_dict()}
@classmethod
def from_dict(cls, payload: dict[str, Any]) -> WorkflowAgent.RequestInfoFunctionArgs:
return cls(request_id=payload.get("request_id", ""), data=payload.get("data"))
if "request_id" not in payload or "request_event" not in payload:
raise ValueError(
"Invalid payload for RequestInfoFunctionArgs. 'request_id' and 'request_event' are required."
)
if not payload["request_id"]:
raise ValueError("request_id cannot be empty.")
@classmethod
def from_json(cls, raw: str) -> WorkflowAgent.RequestInfoFunctionArgs:
try:
parsed: Any = json.loads(raw)
except json.JSONDecodeError as exc:
raise ValueError(f"RequestInfoFunctionArgs JSON payload is malformed: {exc}") from exc
if not isinstance(parsed, dict):
raise ValueError("RequestInfoFunctionArgs JSON payload must decode to a mapping")
return cls.from_dict(cast(dict[str, Any], parsed))
return cls(
request_id=payload.get("request_id", ""),
request_event=WorkflowEvent.from_dict(payload.get("request_event", {})),
)
def __init__(
self,
@@ -129,16 +125,11 @@ class WorkflowAgent(BaseAgent):
**kwargs,
)
self._workflow: Workflow = workflow
self._pending_requests: dict[str, WorkflowEvent[Any]] = {}
@property
def workflow(self) -> Workflow:
return self._workflow
@property
def pending_requests(self) -> dict[str, WorkflowEvent[Any]]:
return self._pending_requests
# region Run Methods
@overload
@@ -182,7 +173,7 @@ class WorkflowAgent(BaseAgent):
Args:
messages: The message(s) to send to the workflow. Required for new runs,
should be None when resuming from checkpoint.
could be None if only restoring the underlying workflow from a checkpoint.
Keyword Args:
stream: If True, returns an async iterable of updates. If False (default),
@@ -416,101 +407,79 @@ class WorkflowAgent(BaseAgent):
Yields:
WorkflowEvent objects from the workflow execution.
"""
# Determine the execution mode based on state.
# The streaming flag controls the workflow's internal streaming mode,
# which affects executor behavior (e.g. AgentExecutor emits different event
# types in streaming vs non-streaming mode).
if bool(self.pending_requests):
function_responses = self._process_pending_requests(input_messages)
# Restore the workflow state if a checkpoint is provided
if checkpoint_id is not None:
if checkpoint_storage is None:
raise AgentInvalidRequestException("checkpoint_storage must be provided when checkpoint_id is provided")
logger.debug(f"Restoring workflow from checkpoint {checkpoint_id}")
# Restore the workflow from checkpoint
if streaming:
async for event in self.workflow.run(
responses=function_responses,
stream=True,
function_invocation_kwargs=function_invocation_kwargs,
client_kwargs=client_kwargs,
):
yield event
else:
for event in await self.workflow.run(
responses=function_responses,
function_invocation_kwargs=function_invocation_kwargs,
client_kwargs=client_kwargs,
):
yield event
elif checkpoint_id is not None:
# Restore the prior workflow state from the checkpoint. Shared
# state (e.g. accumulated conversation history maintained by the
# workflow's executors) survives across turns because Workflow.run
# no longer wipes state per call. Callers who want to deliver a
# new user message after restore should make a second
# `workflow.run(message=...)` call - they are NOT mutually
# exclusive on the same instance, but each must be its own call.
if streaming:
async for event in self.workflow.run(
async for _ in self.workflow.run(
stream=True,
checkpoint_id=checkpoint_id,
checkpoint_storage=checkpoint_storage,
):
pass
else:
_ = await self.workflow.run(
checkpoint_id=checkpoint_id,
checkpoint_storage=checkpoint_storage,
)
if not input_messages:
logger.info("No input messages provided; the workflow has been restored to the checkpoint state.")
return
final_state = self._workflow.status
logger.debug(f"Workflow state: {final_state}")
if final_state == WorkflowRunState.IDLE_WITH_PENDING_REQUESTS:
# Extract function responses from input messages, and ensure that
# only function responses are present in messages if there is any
# pending request.
# NOTE: It is possible that some pending requests are not fulfilled,
# and we will let the workflow to handle this -- the agent does not
# have an opinion on this.
function_responses = self._extract_function_responses(input_messages)
if streaming:
async for event in self.workflow.run(
responses=function_responses,
stream=True,
checkpoint_storage=checkpoint_storage,
function_invocation_kwargs=function_invocation_kwargs,
client_kwargs=client_kwargs,
):
yield event
else:
for event in await self.workflow.run(
checkpoint_id=checkpoint_id,
responses=function_responses,
checkpoint_storage=checkpoint_storage,
function_invocation_kwargs=function_invocation_kwargs,
client_kwargs=client_kwargs,
):
yield event
elif final_state == WorkflowRunState.IDLE:
if streaming:
async for event in self.workflow.run(
message=input_messages,
stream=True,
checkpoint_storage=checkpoint_storage,
function_invocation_kwargs=function_invocation_kwargs,
client_kwargs=client_kwargs,
):
yield event
else:
for event in await self.workflow.run(
message=input_messages,
checkpoint_storage=checkpoint_storage,
function_invocation_kwargs=function_invocation_kwargs,
client_kwargs=client_kwargs,
):
yield event
else:
if streaming:
async for event in self.workflow.run(
message=input_messages,
stream=True,
checkpoint_storage=checkpoint_storage,
function_invocation_kwargs=function_invocation_kwargs,
client_kwargs=client_kwargs,
):
yield event
else:
for event in await self.workflow.run(
message=input_messages,
checkpoint_storage=checkpoint_storage,
function_invocation_kwargs=function_invocation_kwargs,
client_kwargs=client_kwargs,
):
yield event
raise AgentException(f"The underlying workflow is in an invalid state to restart: {final_state}.")
# endregion Run Methods
def _process_pending_requests(self, input_messages: Sequence[Message]) -> dict[str, Any]:
"""Process pending requests by extracting function responses and updating state.
Args:
input_messages: Input messages that may contain function responses.
Returns:
A dictionary mapping request IDs to their response data.
"""
logger.info(f"Continuing workflow to address {len(self.pending_requests)} requests")
# Extract function responses from input messages, and ensure that
# only function responses are present in messages if there is any
# pending request.
function_responses = self._extract_function_responses(input_messages)
# Pop pending requests if fulfilled.
for request_id in list(self.pending_requests.keys()):
if request_id in function_responses:
self.pending_requests.pop(request_id)
# NOTE: It is possible that some pending requests are not fulfilled,
# and we will let the workflow to handle this -- the agent does not
# have an opinion on this.
return function_responses
def _convert_workflow_events_to_agent_response(
self,
response_id: str,
@@ -528,10 +497,10 @@ class WorkflowAgent(BaseAgent):
for output_event in output_events:
if output_event.type == "request_info":
function_call, approval_request = self._process_request_info_event(output_event)
request_content = self._process_request_info_event(output_event)
messages.append(
Message(
contents=[function_call, approval_request],
contents=[request_content],
role="assistant",
author_name=output_event.source_executor_id,
message_id=str(uuid.uuid4()),
@@ -598,38 +567,6 @@ class WorkflowAgent(BaseAgent):
raw_representation=raw_representations,
)
def _process_request_info_event(
self,
event: WorkflowEvent[Any],
) -> tuple[Content, Content]:
"""Convert a request_info event to FunctionCallContent and FunctionApprovalRequestContent.
Args:
event: A WorkflowEvent with type='request_info'.
Returns:
A tuple of (FunctionCallContent, FunctionApprovalRequestContent).
"""
request_id = event.request_id
if not request_id:
raise ValueError("request_info event must have a request_id")
self.pending_requests[request_id] = event
args = self.RequestInfoFunctionArgs(request_id=request_id, data=event.data).to_dict()
function_call = Content.from_function_call(
call_id=request_id,
name=self.REQUEST_INFO_FUNCTION_NAME,
arguments=args,
)
approval_request = Content.from_function_approval_request(
id=request_id,
function_call=function_call,
additional_properties={"request_id": request_id},
)
return function_call, approval_request
def _convert_workflow_event_to_agent_response_updates(
self,
response_id: str,
@@ -731,85 +668,72 @@ class WorkflowAgent(BaseAgent):
]
if event.type == "request_info":
# Store the pending request for later correlation
request_id = event.request_id
if not request_id:
raise ValueError("request_info event must have a request_id")
self.pending_requests[request_id] = event
args = self.RequestInfoFunctionArgs(request_id=request_id, data=event.data).to_dict()
function_call = Content.from_function_call(
call_id=request_id,
name=self.REQUEST_INFO_FUNCTION_NAME,
arguments=args,
)
approval_request = Content.from_function_approval_request(
id=request_id,
function_call=function_call,
additional_properties={"request_id": request_id},
)
request_content = self._process_request_info_event(event)
return [
AgentResponseUpdate(
contents=[function_call, approval_request],
contents=[request_content],
role="assistant",
author_name=self.name,
response_id=response_id,
message_id=str(uuid.uuid4()),
created_at=datetime.now(tz=timezone.utc).strftime("%Y-%m-%dT%H:%M:%S.%fZ"),
raw_representation=event,
)
]
# Ignore workflow-internal events
return []
def _process_request_info_event(
self,
event: WorkflowEvent[Any],
) -> Content:
"""Convert a request_info event to FunctionApprovalRequestContent.
Args:
event: A WorkflowEvent with type='request_info'.
Returns:
A content object representing the request info. The content can be a `function_approval_request`
or a `function_call` depending on the structure of the event data.
Note:
If the event data is already a FunctionApprovalRequestContent, it will be returned as-is.
"""
if isinstance(event.data, Content) and event.data.user_input_request:
# Return the event data as-is if it's already a properly formed FunctionApprovalRequestContent
return event.data
request_id = event.request_id
args = self.RequestInfoFunctionArgs(request_id=request_id, request_event=event).to_dict()
return Content.from_function_call(
call_id=request_id,
name=self.REQUEST_INFO_FUNCTION_NAME,
arguments=args,
)
def _extract_function_responses(self, input_messages: Sequence[Message]) -> dict[str, Any]:
"""Extract function responses from input messages."""
"""Extract function responses from input messages.
The responses are for pending requests that the workflow is waiting on, and
will be passed to the workflow. The pending requests are processed to either
`function_approval_request` or `function_call` content by `_process_request_info_event`.
"""
function_responses: dict[str, Any] = {}
for message in input_messages:
for content in message.contents:
if content.type == "function_approval_response":
# Parse the function arguments to recover request payload
arguments_payload = content.function_call.arguments # type: ignore[attr-defined, union-attr]
if isinstance(arguments_payload, str):
try:
parsed_args = self.RequestInfoFunctionArgs.from_json(arguments_payload)
except ValueError as exc:
raise AgentInvalidResponseException(
"FunctionApprovalResponseContent arguments must decode to a mapping."
) from exc
elif isinstance(arguments_payload, dict):
parsed_args = self.RequestInfoFunctionArgs.from_dict(arguments_payload)
else:
raise AgentInvalidResponseException(
"FunctionApprovalResponseContent arguments must be a mapping or JSON string."
)
request_id = parsed_args.request_id or content.id # type: ignore[attr-defined]
if not content.approved: # type: ignore[attr-defined]
raise AgentInvalidResponseException(f"Request '{request_id}' was not approved by the caller.")
if request_id in self.pending_requests:
function_responses[request_id] = parsed_args.data
elif bool(self.pending_requests):
raise AgentInvalidRequestException(
"Only responses for pending requests are allowed when there are outstanding approvals."
)
request_id: str = content.id # type: ignore[assignment]
function_responses[request_id] = content
elif content.type == "function_result":
request_id = content.call_id # type: ignore[attr-defined]
if request_id in self.pending_requests:
response_data = content.result if hasattr(content, "result") else str(content) # type: ignore[attr-defined]
function_responses[request_id] = response_data
elif bool(self.pending_requests):
raise AgentInvalidRequestException(
"Only function responses for pending requests are allowed while requests are outstanding."
)
response_data = content.result if hasattr(content, "result") else str(content) # type: ignore[attr-defined]
function_responses[content.call_id] = response_data # type: ignore
else:
if bool(self.pending_requests):
raise AgentInvalidResponseException(
"Unexpected content type while awaiting request info responses."
)
raise AgentInvalidResponseException(
"Unexpected content type while awaiting request info responses."
)
return function_responses
def _extract_contents(self, data: Any) -> list[Content]:
@@ -429,15 +429,30 @@ class AgentExecutor(Executor):
function_invocation_kwargs=function_invocation_kwargs,
client_kwargs=client_kwargs,
)
await ctx.yield_output(response)
# Handle any user input requests
if response.user_input_requests:
user_input_request_count = len(response.user_input_requests)
total_message_content_count = sum(len(msg.contents) for msg in response.messages)
if user_input_request_count != total_message_content_count:
logger.warning(
"Response %s contains %d user input requests but total message contents are %d. "
"This indicates the response contains both user input requests and message contents. "
"Double check if this is the intended behavior, as non user input request contents in "
"this response will not be emitted.",
response.response_id,
user_input_request_count,
total_message_content_count,
)
for user_input_request in response.user_input_requests:
self._pending_agent_requests[user_input_request.id] = user_input_request # type: ignore[index]
await ctx.request_info(user_input_request, Content)
await ctx.request_info(user_input_request, Content, request_id=user_input_request.id)
return None
# Only yield output if the response is complete and not waiting for user input.
# This is to avoid emitting two events of different types ('output' and 'request_info')
# that carry the same payload.
await ctx.yield_output(response)
return response
async def _run_agent_streaming(self, ctx: WorkflowContext[Never, AgentResponseUpdate]) -> AgentResponse | None:
@@ -472,9 +487,25 @@ class AgentExecutor(Executor):
)
async for update in stream:
updates.append(update)
await ctx.yield_output(update)
if update.user_input_requests:
user_input_request_count = len(update.user_input_requests)
total_message_content_count = len(update.contents)
if user_input_request_count != total_message_content_count:
logger.warning(
"Response update %s contains %d user input requests but total message contents are %d. "
"This indicates the response update contains both user input requests and message contents. "
"Double check if this is the intended behavior, as non user input request contents will "
"not be emitted.",
update.response_id,
user_input_request_count,
total_message_content_count,
)
streamed_user_input_requests.extend(update.user_input_requests)
else:
# Only yield output events for updates that do not contain user input requests.
# This is to avoid emitting two events of different types ('output' and 'request_info')
# that carry the same payload.
await ctx.yield_output(update)
# Prefer stream finalization when available so result hooks run
# (e.g., thread conversation updates). Fall back to reconstructing from updates
@@ -509,7 +540,7 @@ class AgentExecutor(Executor):
if user_input_requests:
for user_input_request in user_input_requests:
self._pending_agent_requests[user_input_request.id] = user_input_request # type: ignore[index]
await ctx.request_info(user_input_request, Content)
await ctx.request_info(user_input_request, Content, request_id=user_input_request.id)
return None
return response
@@ -360,6 +360,22 @@ class Workflow(DictConvertible):
# Flag to prevent concurrent workflow executions
self._is_running = False
# Current run-level status of this workflow instance. Updated in lockstep with
# the status events emitted from `_run_workflow_with_tracing`. Defaults to IDLE
# for a freshly built workflow that has not yet been run.
self._status: WorkflowRunState = WorkflowRunState.IDLE
@property
def status(self) -> WorkflowRunState:
"""Return the current run-level status of this workflow instance.
Mirrors the most recent status event emitted by the workflow. Safe to read at
any time: workflows run on a single asyncio event loop, and the underlying
attribute is a single enum reference whose assignment is atomic under the
CPython GIL, so no locking is required.
"""
return self._status
def _ensure_not_running(self) -> None:
"""Ensure the workflow is not already running."""
if self._is_running:
@@ -513,8 +529,9 @@ class Workflow(DictConvertible):
with _framework_event_origin():
started = WorkflowEvent.started()
yield started # noqa: RUF070
self._status = WorkflowRunState.IN_PROGRESS
with _framework_event_origin():
in_progress = WorkflowEvent.status(WorkflowRunState.IN_PROGRESS)
in_progress = WorkflowEvent.status(self._status)
yield in_progress # noqa: RUF070
# Per-run reset for fresh-message runs only. We deliberately
@@ -569,17 +586,20 @@ class Workflow(DictConvertible):
if event.type == "request_info" and not emitted_in_progress_pending:
emitted_in_progress_pending = True
self._status = WorkflowRunState.IN_PROGRESS_PENDING_REQUESTS
with _framework_event_origin():
pending_status = WorkflowEvent.status(WorkflowRunState.IN_PROGRESS_PENDING_REQUESTS)
pending_status = WorkflowEvent.status(self._status)
yield pending_status # noqa: RUF070
# Workflow runs until idle - emit final status based on whether requests are pending
if saw_request:
self._status = WorkflowRunState.IDLE_WITH_PENDING_REQUESTS
with _framework_event_origin():
terminal_status = WorkflowEvent.status(WorkflowRunState.IDLE_WITH_PENDING_REQUESTS)
terminal_status = WorkflowEvent.status(self._status)
yield terminal_status
else:
self._status = WorkflowRunState.IDLE
with _framework_event_origin():
terminal_status = WorkflowEvent.status(WorkflowRunState.IDLE)
terminal_status = WorkflowEvent.status(self._status)
yield terminal_status
span.add_event(OtelAttr.WORKFLOW_COMPLETED)
@@ -593,6 +613,7 @@ class Workflow(DictConvertible):
with _framework_event_origin():
failed_event = WorkflowEvent.failed(details)
yield failed_event # noqa: RUF070
self._status = WorkflowRunState.FAILED
with _framework_event_origin():
failed_status = WorkflowEvent.status(WorkflowRunState.FAILED)
yield failed_status # noqa: RUF070
@@ -80,6 +80,7 @@ __all__ = [
"EmbeddingTelemetryLayer",
"OtelAttr",
"configure_otel_providers",
"create_mcp_client_span",
"create_metric_views",
"create_resource",
"disable_instrumentation",
@@ -87,6 +88,7 @@ __all__ = [
"enable_sensitive_telemetry",
"get_meter",
"get_tracer",
"set_mcp_span_error",
]
@@ -110,7 +112,6 @@ INNER_ACCUMULATED_USAGE: Final[contextvars.ContextVar[UsageDetails | None]] = co
"inner_accumulated_usage", default=None
)
OTEL_METRICS: Final[str] = "__otel_metrics__"
TOKEN_USAGE_BUCKET_BOUNDARIES: Final[tuple[float, ...]] = (
1,
@@ -292,6 +293,14 @@ class OtelAttr(str, Enum):
AGENT_CREATE_OPERATION = "create_agent"
AGENT_INVOKE_OPERATION = "invoke_agent"
# MCP attributes (https://opentelemetry.io/docs/specs/semconv/gen-ai/mcp/)
MCP_METHOD_NAME = "mcp.method.name"
MCP_PROTOCOL_VERSION = "mcp.protocol.version"
MCP_SESSION_ID = "mcp.session.id"
PROMPT_NAME = "gen_ai.prompt.name"
NETWORK_TRANSPORT = "network.transport"
NETWORK_PROTOCOL_NAME = "network.protocol.name"
# Agent Framework specific attributes
MEASUREMENT_FUNCTION_TAG_NAME = "agent_framework.function.name"
MEASUREMENT_FUNCTION_INVOCATION_DURATION = "agent_framework.function.invocation.duration"
@@ -2013,6 +2022,61 @@ def get_function_span(
)
# region MCP span helpers
@contextlib.contextmanager
def create_mcp_client_span(
method_name: str,
target: str | None = None,
attributes: dict[str, Any] | None = None,
) -> Generator[trace.Span, Any, Any]:
"""Create an MCP client span per OTel MCP semantic conventions.
Span name follows the format ``{mcp.method.name} {target}`` when a target
is available, otherwise just ``{mcp.method.name}``.
See: https://opentelemetry.io/docs/specs/semconv/gen-ai/mcp/#client
Args:
method_name: The MCP method name (e.g. ``initialize``, ``tools/call``).
target: Optional low-cardinality target (tool name, prompt name).
attributes: Additional span attributes.
"""
span_name = f"{method_name} {target}" if target else method_name
attrs: dict[str, Any] = {OtelAttr.MCP_METHOD_NAME: method_name}
if attributes:
attrs.update(attributes)
tracer = get_tracer() if OBSERVABILITY_SETTINGS.ENABLED else trace.NoOpTracer()
span = tracer.start_span(span_name, kind=trace.SpanKind.CLIENT, attributes=attrs)
with trace.use_span(
span=span,
end_on_exit=True,
record_exception=True,
set_status_on_exception=True,
) as current_span:
yield current_span
def set_mcp_span_error(
span: trace.Span,
error_type: str,
description: str | None = None,
) -> None:
"""Set error status and ``error.type`` on an MCP span.
Args:
span: The span to mark as errored.
error_type: The error type string (e.g. ``tool_error``, exception class name).
description: Optional description (e.g. JSON-RPC error message).
"""
span.set_attribute(OtelAttr.ERROR_TYPE, error_type)
span.set_status(trace.StatusCode.ERROR, description=description)
# endregion
@contextlib.contextmanager
def _activate_span(span: trace.Span) -> Generator[None]:
"""Attach ``span`` as the current span in the OpenTelemetry context.
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,376 @@
# Copyright (c) Microsoft. All rights reserved.
"""Tests for MCP client span instrumentation per OTel GenAI Semantic Conventions.
See: https://opentelemetry.io/docs/specs/semconv/gen-ai/mcp/#client
"""
from __future__ import annotations
from typing import Any
from unittest.mock import AsyncMock, Mock
import pytest
from mcp import types
from mcp.shared.exceptions import McpError
from mcp.types import ErrorData
from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter
from opentelemetry.trace import SpanKind, StatusCode
from agent_framework import MCPStdioTool, MCPStreamableHTTPTool, MCPWebsocketTool
from agent_framework._mcp import MCPTool
from agent_framework.exceptions import ToolExecutionException
from agent_framework.observability import OtelAttr
# region helpers
def _make_connected_mcp_tool(
name: str = "test-mcp",
*,
supports_tools: bool = True,
supports_prompts: bool = True,
) -> MCPTool:
"""Create an MCPTool with a mocked session, ready for testing."""
tool = MCPTool(name=name)
tool.session = AsyncMock()
tool.is_connected = True
tool._supports_tools = supports_tools
tool._supports_prompts = supports_prompts
tool.load_tools_flag = True
tool.load_prompts_flag = True
return tool
def _make_tool_list_result(
tools: list[dict[str, Any]] | None = None,
) -> Mock:
"""Create a mock ListToolsResult."""
if tools is None:
tools = [{"name": "get-weather", "description": "Get weather", "inputSchema": {"type": "object"}}]
result = Mock()
result.tools = [
types.Tool(name=t["name"], description=t.get("description", ""), inputSchema=t.get("inputSchema", {}))
for t in tools
]
result.nextCursor = None
return result
def _make_prompt_list_result(
prompts: list[dict[str, Any]] | None = None,
) -> Mock:
"""Create a mock ListPromptsResult."""
if prompts is None:
prompts = [{"name": "analyze-code", "description": "Analyze code"}]
result = Mock()
result.prompts = [
types.Prompt(name=p["name"], description=p.get("description", ""), arguments=None) for p in prompts
]
result.nextCursor = None
return result
def _make_call_tool_result(text: str = "result", is_error: bool = False) -> Mock:
"""Create a mock CallToolResult."""
result = Mock()
result.isError = is_error
result.content = [types.TextContent(type="text", text=text)]
return result
def _make_get_prompt_result(text: str = "prompt result") -> types.GetPromptResult:
"""Create a mock GetPromptResult."""
return types.GetPromptResult(
description="test prompt",
messages=[
types.PromptMessage(
role="user",
content=types.TextContent(type="text", text=text),
)
],
)
# endregion
# region initialize span
async def test_mcp_initialize_span(span_exporter: InMemorySpanExporter):
"""session.initialize() should produce an MCP CLIENT span named 'initialize'."""
tool = MCPTool(name="test-server")
mock_session_cls = AsyncMock()
init_result = Mock()
init_result.capabilities = None
init_result.protocolVersion = "2025-06-18"
mock_session_cls.initialize = AsyncMock(return_value=init_result)
# Create a mock transport context manager
mock_transport = AsyncMock()
mock_transport.__aenter__ = AsyncMock(return_value=(Mock(), Mock()))
mock_transport.__aexit__ = AsyncMock(return_value=False)
# Mock get_mcp_client and the session creation
tool.session = None
tool.load_tools_flag = False
tool.load_prompts_flag = False
span_exporter.clear()
with pytest.MonkeyPatch.context() as m:
m.setattr(tool, "get_mcp_client", lambda: mock_transport)
async def patched_connect(self_: Any, *, reset: bool = False, load_configured: bool = True) -> None:
# Simulate _connect_on_owner: create initialize span and call session.initialize()
from agent_framework._mcp import create_mcp_client_span
from agent_framework.observability import OtelAttr
with create_mcp_client_span("initialize", attributes=self_._mcp_base_span_attributes()) as init_span:
result = await mock_session_cls.initialize()
protocol_version = getattr(result, "protocolVersion", None)
if protocol_version:
init_span.set_attribute(OtelAttr.MCP_PROTOCOL_VERSION, protocol_version)
self_.session = mock_session_cls
self_.is_connected = True
m.setattr(MCPTool, "_connect_on_owner", patched_connect)
await tool.connect()
mock_session_cls.initialize.assert_awaited_once()
spans = span_exporter.get_finished_spans()
init_spans = [s for s in spans if s.name == "initialize"]
assert len(init_spans) == 1
span = init_spans[0]
assert span.kind == SpanKind.CLIENT
assert span.attributes[OtelAttr.MCP_METHOD_NAME] == "initialize"
assert span.attributes.get(OtelAttr.MCP_PROTOCOL_VERSION) == "2025-06-18"
# endregion
# region tools/list span
async def test_mcp_tools_list_span(span_exporter: InMemorySpanExporter):
"""session.list_tools() should produce an MCP CLIENT span named 'tools/list'."""
tool = _make_connected_mcp_tool()
tool.session.list_tools = AsyncMock(return_value=_make_tool_list_result())
span_exporter.clear()
await tool.load_tools()
spans = span_exporter.get_finished_spans()
list_spans = [s for s in spans if s.name == "tools/list"]
assert len(list_spans) == 1
span = list_spans[0]
assert span.kind == SpanKind.CLIENT
assert span.attributes[OtelAttr.MCP_METHOD_NAME] == "tools/list"
# endregion
# region prompts/list span
async def test_mcp_prompts_list_span(span_exporter: InMemorySpanExporter):
"""session.list_prompts() should produce an MCP CLIENT span named 'prompts/list'."""
tool = _make_connected_mcp_tool()
tool.session.list_prompts = AsyncMock(return_value=_make_prompt_list_result())
span_exporter.clear()
await tool.load_prompts()
spans = span_exporter.get_finished_spans()
list_spans = [s for s in spans if s.name == "prompts/list"]
assert len(list_spans) == 1
span = list_spans[0]
assert span.kind == SpanKind.CLIENT
assert span.attributes[OtelAttr.MCP_METHOD_NAME] == "prompts/list"
# endregion
# region tools/call span
async def test_mcp_tools_call_creates_client_span_when_no_parent(span_exporter: InMemorySpanExporter):
"""Direct call_tool() without FunctionTool wrapper creates new MCP CLIENT span."""
tool = _make_connected_mcp_tool()
tool.session.call_tool = AsyncMock(return_value=_make_call_tool_result("hello"))
span_exporter.clear()
result = await tool.call_tool("get-weather", city="Seattle")
assert result is not None
spans = span_exporter.get_finished_spans()
call_spans = [s for s in spans if "tools/call" in s.name]
assert len(call_spans) == 1
span = call_spans[0]
assert span.kind == SpanKind.CLIENT
assert span.name == "tools/call get-weather"
assert span.attributes[OtelAttr.MCP_METHOD_NAME] == "tools/call"
assert span.attributes[OtelAttr.TOOL_NAME] == "get-weather"
async def test_mcp_tools_call_tool_error_sets_error_type(span_exporter: InMemorySpanExporter):
"""When CallToolResult.isError is true, error.type should be 'tool_error' per MCP spec."""
tool = _make_connected_mcp_tool()
tool.session.call_tool = AsyncMock(return_value=_make_call_tool_result("bad input", is_error=True))
span_exporter.clear()
with pytest.raises(ToolExecutionException):
await tool.call_tool("get-weather", city="invalid")
spans = span_exporter.get_finished_spans()
call_spans = [s for s in spans if "tools/call" in s.name]
assert len(call_spans) == 1
span = call_spans[0]
assert span.attributes.get(OtelAttr.ERROR_TYPE) == "tool_error"
assert span.status.status_code == StatusCode.ERROR
async def test_mcp_tools_call_mcp_error_sets_error_type(span_exporter: InMemorySpanExporter):
"""When session.call_tool() raises McpError, error.type should be the exception class name."""
tool = _make_connected_mcp_tool()
tool.session.call_tool = AsyncMock(side_effect=McpError(ErrorData(code=-32600, message="invalid request")))
span_exporter.clear()
with pytest.raises(ToolExecutionException):
await tool.call_tool("get-weather")
spans = span_exporter.get_finished_spans()
call_spans = [s for s in spans if "tools/call" in s.name]
assert len(call_spans) == 1
span = call_spans[0]
assert span.attributes.get(OtelAttr.ERROR_TYPE) == "McpError"
assert span.status.status_code == StatusCode.ERROR
# endregion
# region prompts/get span
async def test_mcp_prompts_get_creates_client_span(span_exporter: InMemorySpanExporter):
"""get_prompt() should always create a new MCP CLIENT span (not enrich execute_tool)."""
tool = _make_connected_mcp_tool()
tool.session.get_prompt = AsyncMock(return_value=_make_get_prompt_result("code analysis"))
span_exporter.clear()
result = await tool.get_prompt("analyze-code", language="python")
assert "code analysis" in result
spans = span_exporter.get_finished_spans()
prompt_spans = [s for s in spans if "prompts/get" in s.name]
assert len(prompt_spans) == 1
span = prompt_spans[0]
assert span.kind == SpanKind.CLIENT
assert span.name == "prompts/get analyze-code"
assert span.attributes[OtelAttr.MCP_METHOD_NAME] == "prompts/get"
assert span.attributes[OtelAttr.PROMPT_NAME] == "analyze-code"
async def test_mcp_prompts_get_mcp_error_sets_error_type(span_exporter: InMemorySpanExporter):
"""When session.get_prompt() raises McpError, the span should have error.type and ERROR status."""
tool = _make_connected_mcp_tool()
tool.session.get_prompt = AsyncMock(
side_effect=McpError(ErrorData(code=-32602, message="prompt not found"))
)
span_exporter.clear()
with pytest.raises(ToolExecutionException):
await tool.get_prompt("missing-prompt")
spans = span_exporter.get_finished_spans()
prompt_spans = [s for s in spans if "prompts/get" in s.name]
assert len(prompt_spans) == 1
span = prompt_spans[0]
assert span.attributes.get(OtelAttr.ERROR_TYPE) == "McpError"
assert span.status.status_code == StatusCode.ERROR
# endregion
# region transport attributes
def test_mcp_stdio_tool_transport_attributes():
"""MCPStdioTool should have network.transport='pipe'."""
tool = MCPStdioTool(name="test", command="python")
attrs = tool._mcp_base_span_attributes()
assert attrs[OtelAttr.NETWORK_TRANSPORT] == "pipe"
assert OtelAttr.ADDRESS not in attrs
def test_mcp_http_tool_transport_attributes():
"""MCPStreamableHTTPTool should have tcp transport and URL-based server address/port."""
tool = MCPStreamableHTTPTool(name="test", url="https://api.example.com:8443/mcp")
attrs = tool._mcp_base_span_attributes()
assert attrs[OtelAttr.NETWORK_TRANSPORT] == "tcp"
assert attrs[OtelAttr.NETWORK_PROTOCOL_NAME] == "http"
assert attrs[OtelAttr.ADDRESS] == "api.example.com"
assert attrs[OtelAttr.PORT] == 8443
def test_mcp_http_tool_default_port():
"""MCPStreamableHTTPTool should default to 443 for https."""
tool = MCPStreamableHTTPTool(name="test", url="https://api.example.com/mcp")
attrs = tool._mcp_base_span_attributes()
assert attrs[OtelAttr.PORT] == 443
def test_mcp_http_tool_http_default_port():
"""MCPStreamableHTTPTool should default to 80 for http."""
tool = MCPStreamableHTTPTool(name="test", url="http://localhost/mcp")
attrs = tool._mcp_base_span_attributes()
assert attrs[OtelAttr.PORT] == 80
def test_mcp_websocket_tool_transport_attributes():
"""MCPWebsocketTool should have tcp transport and URL-based server address/port."""
tool = MCPWebsocketTool(name="test", url="wss://ws.example.com:9090/mcp")
attrs = tool._mcp_base_span_attributes()
assert attrs[OtelAttr.NETWORK_TRANSPORT] == "tcp"
assert attrs[OtelAttr.NETWORK_PROTOCOL_NAME] == "websocket"
assert attrs[OtelAttr.ADDRESS] == "ws.example.com"
assert attrs[OtelAttr.PORT] == 9090
def test_mcp_websocket_tool_default_port():
"""MCPWebsocketTool should default to 443 for wss."""
tool = MCPWebsocketTool(name="test", url="wss://ws.example.com/mcp")
attrs = tool._mcp_base_span_attributes()
assert attrs[OtelAttr.PORT] == 443
# endregion
# region observability disabled
@pytest.mark.parametrize("enable_instrumentation", [False], indirect=True)
async def test_mcp_spans_not_created_when_observability_disabled(span_exporter: InMemorySpanExporter):
"""No MCP spans should be created when observability is disabled."""
tool = _make_connected_mcp_tool()
tool.session.list_tools = AsyncMock(return_value=_make_tool_list_result())
tool.session.call_tool = AsyncMock(return_value=_make_call_tool_result("ok"))
span_exporter.clear()
await tool.load_tools()
await tool.call_tool("get-weather", city="Seattle")
spans = span_exporter.get_finished_spans()
assert len(spans) == 0
# endregion
@@ -25,6 +25,7 @@ from agent_framework import (
prepend_agent_framework_to_user_agent,
tool,
)
from agent_framework._serialization import make_json_safe
from agent_framework.observability import (
ROLE_EVENT_MAP,
AgentTelemetryLayer,
@@ -3195,17 +3196,15 @@ def test_capture_messages_with_prepared_request_info_function_call_arguments(spa
from opentelemetry import trace
from agent_framework import WorkflowAgent
@dataclasses.dataclass
class HandoffRequest:
target_agent: str
reason: str
arguments = WorkflowAgent.RequestInfoFunctionArgs(
request_id="call_dc",
data=HandoffRequest(target_agent="helper", reason="overflow"),
).to_dict()
arguments = {
"request_id": "call_dc",
"data": make_json_safe(HandoffRequest(target_agent="helper", reason="overflow")),
}
msg = Message(
role="assistant",
contents=[
@@ -699,3 +699,171 @@ async def test_resolve_executor_kwargs_empty_per_executor_does_not_fallback_to_g
resolved = {"exec_a": {}, GLOBAL_KWARGS_KEY: {"global_key": "global_val"}}
result = executor._resolve_executor_kwargs(resolved) # pyright: ignore[reportPrivateUsage]
assert result == {}
# region Tool approval emission
class _ApprovalEmittingAgent(BaseAgent):
"""Agent that returns a single ``function_approval_request`` Content.
Used to verify that ``AgentExecutor`` does *not* surface the approval
payload via both an ``output`` event and a ``request_info`` event in the
same superstep — only the ``request_info`` event must carry it.
"""
def __init__(
self,
*,
approval_request_id: str = "apr_1",
tool_name: str = "delete_file",
tool_arguments: dict[str, Any] | None = None,
**kwargs: Any,
):
super().__init__(**kwargs)
self._approval_request_id = approval_request_id
self._tool_name = tool_name
self._tool_arguments: dict[str, Any] = tool_arguments or {"path": "/tmp/secret.txt"}
self.run_count = 0
def _build_approval_content(self) -> Content:
function_call = Content.from_function_call(
call_id=self._approval_request_id,
name=self._tool_name,
arguments=self._tool_arguments,
)
return Content.from_function_approval_request(id=self._approval_request_id, function_call=function_call)
@overload
def run(
self,
messages: AgentRunInputs | None = ...,
*,
stream: Literal[False] = ...,
session: AgentSession | None = ...,
**kwargs: Any,
) -> Awaitable[AgentResponse[Any]]: ...
@overload
def run(
self,
messages: AgentRunInputs | None = ...,
*,
stream: Literal[True],
session: AgentSession | None = ...,
**kwargs: Any,
) -> ResponseStream[AgentResponseUpdate, AgentResponse[Any]]: ...
def run(
self,
messages: AgentRunInputs | None = None,
*,
stream: bool = False,
session: AgentSession | None = None,
**kwargs: Any,
) -> Awaitable[AgentResponse[Any]] | ResponseStream[AgentResponseUpdate, AgentResponse[Any]]:
self.run_count += 1
approval = self._build_approval_content()
if stream:
async def _stream() -> AsyncIterable[AgentResponseUpdate]:
yield AgentResponseUpdate(contents=[approval], role="assistant")
return ResponseStream(_stream(), finalizer=AgentResponse.from_updates)
async def _run() -> AgentResponse:
return AgentResponse(messages=[Message("assistant", [approval])])
return _run()
def _has_approval_payload(event: WorkflowEvent[Any]) -> bool:
"""Return True if the event's data carries a ``function_approval_request`` content."""
data: Any = event.data
def _contents_of(value: Any) -> list[Content]:
if isinstance(value, AgentResponseUpdate):
return list(value.contents)
if isinstance(value, AgentResponse):
return [c for m in value.messages for c in m.contents]
if isinstance(value, AgentExecutorResponse):
return [c for m in value.agent_response.messages for c in m.contents]
if isinstance(value, Message):
return list(value.contents)
if isinstance(value, Content):
return [value]
return []
return any(c.type == "function_approval_request" for c in _contents_of(data))
async def test_agent_executor_does_not_double_emit_approval_non_streaming() -> None:
"""Non-streaming: approval payload must only appear in the ``request_info`` event.
Regression test for the bug where ``AgentExecutor._run_agent`` first
``yield_output``-ed the response (carrying the approval Content) and then
additionally emitted a ``request_info`` event for the same payload.
"""
agent = _ApprovalEmittingAgent(id="approve_agent", name="ApproveAgent", approval_request_id="apr_ns_1")
executor = AgentExecutor(agent, id="approve_exec")
workflow = WorkflowBuilder(start_executor=executor).build()
request_info_events: list[WorkflowEvent[Any]] = []
output_events: list[WorkflowEvent[Any]] = []
for event in await workflow.run("please delete it"):
if event.type == "request_info":
request_info_events.append(event)
elif event.type == "output":
output_events.append(event)
assert len(request_info_events) == 1
assert _has_approval_payload(request_info_events[0])
# The approval payload must not also be surfaced as a workflow output.
assert not any(_has_approval_payload(e) for e in output_events)
assert agent.run_count == 1
async def test_agent_executor_does_not_double_emit_approval_streaming() -> None:
"""Streaming: per-update approval payload must not be ``yield_output``-ed."""
agent = _ApprovalEmittingAgent(id="approve_agent_s", name="ApproveAgentS", approval_request_id="apr_st_1")
executor = AgentExecutor(agent, id="approve_exec_s")
workflow = WorkflowBuilder(start_executor=executor).build()
request_info_events: list[WorkflowEvent[Any]] = []
output_events: list[WorkflowEvent[Any]] = []
async for event in workflow.run("please delete it", stream=True):
if event.type == "request_info":
request_info_events.append(event)
elif event.type == "output":
output_events.append(event)
assert len(request_info_events) == 1
assert _has_approval_payload(request_info_events[0])
assert not any(_has_approval_payload(e) for e in output_events)
assert agent.run_count == 1
async def test_agent_executor_request_info_uses_user_input_request_id() -> None:
"""``ctx.request_info`` must register the request under the agent's approval id.
This makes the workflow's pending-request id round-trip with the
``function_approval_response.id`` the caller echoes back, so
``Workflow._send_responses_internal`` can look it up directly.
"""
agent = _ApprovalEmittingAgent(id="approve_agent_id", name="ApproveAgentId", approval_request_id="apr_match")
executor = AgentExecutor(agent, id="approve_exec_id")
workflow = WorkflowBuilder(start_executor=executor).build()
request_info_events: list[WorkflowEvent[Any]] = []
async for event in workflow.run("please delete it", stream=True):
if event.type == "request_info":
request_info_events.append(event)
assert len(request_info_events) == 1
assert request_info_events[0].request_id == "apr_match"
# endregion Tool approval emission
@@ -1,6 +1,5 @@
# Copyright (c) Microsoft. All rights reserved.
import json
import uuid
from collections.abc import Awaitable, Sequence
from dataclasses import dataclass
@@ -30,6 +29,20 @@ from agent_framework import (
handler,
response_handler,
)
from agent_framework._workflows._typing_utils import deserialize_type
@dataclass
class HandoffRequest:
"""Module-level dataclass used by request_info tests.
Defined at module scope (not nested inside a test method) so
``serialize_type``/``deserialize_type`` can round-trip the request_type via
the importable qualified name ``tests.workflow.test_workflow_agent.HandoffRequest``.
"""
target_agent: str
reason: str
class SimpleExecutor(Executor):
@@ -240,52 +253,45 @@ class TestWorkflowAgent:
# Should have received an approval request for the request info
assert len(updates) > 0
approval_update: AgentResponseUpdate | None = None
request_update: AgentResponseUpdate | None = None
for update in updates:
if any(content.type == "function_approval_request" for content in update.contents):
approval_update = update
if any(content.type == "function_call" for content in update.contents):
request_update = update
break
assert approval_update is not None, "Should have received a request_info approval request"
assert request_update is not None, "Should have received a request_info wrapped in a function_call content"
function_call = next(content for content in approval_update.contents if content.type == "function_call")
approval_request = next(
content for content in approval_update.contents if content.type == "function_approval_request"
)
request_function_call = next(content for content in request_update.contents if content.type == "function_call")
assert request_function_call.call_id is not None
# Verify the function call has expected structure
assert function_call.call_id is not None
assert function_call.name == "request_info"
assert isinstance(function_call.arguments, dict)
assert function_call.arguments.get("request_id") == approval_request.id
assert request_function_call.name == WorkflowAgent.REQUEST_INFO_FUNCTION_NAME
assert isinstance(request_function_call.arguments, dict)
assert request_function_call.arguments.get("request_id") is not None
assert request_function_call.arguments.get("request_event") is not None
request_event = request_function_call.arguments["request_event"]
assert request_event.get("type") == "request_info"
assert deserialize_type(request_event.get("response_type")) is str
# Approval request should reference the same function call
assert approval_request.id is not None
assert approval_request.function_call is not None
assert approval_request.function_call.call_id == function_call.call_id
assert approval_request.function_call.name == function_call.name
deserialized_args = WorkflowAgent.RequestInfoFunctionArgs.from_dict(request_function_call.arguments)
assert deserialized_args.request_id == request_function_call.call_id
assert isinstance(deserialized_args.request_event, WorkflowEvent)
assert deserialized_args.request_event.type == "request_info"
assert deserialized_args.request_event.data == "Mock request data"
assert deserialized_args.request_event.response_type is str
# Verify the request is tracked in pending_requests
assert len(agent.pending_requests) == 1
assert function_call.call_id in agent.pending_requests
pending_requests = await workflow._runner_context.get_pending_request_info_events()
assert len(pending_requests) == 1
assert request_function_call.call_id in pending_requests
# Now provide an approval response with updated arguments to test continuation
response_args = WorkflowAgent.RequestInfoFunctionArgs(
request_id=approval_request.id,
data="User provided answer",
).to_dict()
approval_response = Content.from_function_approval_response(
approved=True,
id=approval_request.id,
function_call=Content.from_function_call(
call_id=function_call.call_id,
name=function_call.name,
arguments=response_args,
),
# Now provide a function result response with updated arguments to test continuation
function_result = Content.from_function_result(
call_id=request_function_call.call_id,
result="Mock response to request info",
)
response_message = Message(role="user", contents=[approval_response])
response_message = Message(role="user", contents=[function_result])
# Continue the workflow with the response
continuation_result = await agent.run(response_message)
@@ -294,16 +300,11 @@ class TestWorkflowAgent:
assert isinstance(continuation_result, AgentResponse)
# Verify cleanup - pending requests should be cleared after function response handling
assert len(agent.pending_requests) == 0
pending_requests = await workflow._runner_context.get_pending_request_info_events()
assert len(pending_requests) == 0
def test_request_info_dataclass_arguments_are_serialized_when_content_is_created(self) -> None:
"""Test WorkflowAgent prepares request_info arguments before observability captures messages."""
@dataclass
class HandoffRequest:
target_agent: str
reason: str
executor = SimpleExecutor(id="executor1", response_text="Response")
workflow = WorkflowBuilder(start_executor=executor).build()
agent = WorkflowAgent(workflow=workflow, name="Request Test Agent")
@@ -314,14 +315,367 @@ class TestWorkflowAgent:
response_type=str,
)
function_call, approval_request = agent._process_request_info_event(event) # pyright: ignore[reportPrivateUsage]
request_function_call = agent._process_request_info_event(event) # pyright: ignore[reportPrivateUsage]
assert function_call.arguments == {
"request_id": "request_123",
"data": {"target_agent": "helper", "reason": "overflow"},
}
assert approval_request.function_call is function_call
assert json.loads(json.dumps(function_call.arguments)) == function_call.arguments
assert request_function_call.call_id == "request_123"
assert isinstance(request_function_call.arguments, dict)
assert request_function_call.arguments.get("request_event") is not None
request_event = request_function_call.arguments["request_event"]
assert request_event.get("type") == "request_info"
assert request_event.get("request_id") == "request_123"
assert request_event.get("source_executor_id") == "executor1"
assert deserialize_type(request_event.get("response_type")) is str
assert request_event.get("data") == HandoffRequest(target_agent="helper", reason="overflow")
deserialized_args = WorkflowAgent.RequestInfoFunctionArgs.from_dict(request_function_call.arguments)
assert deserialized_args.request_id == "request_123"
assert isinstance(deserialized_args.request_event, WorkflowEvent)
assert deserialized_args.request_event.type == "request_info"
assert deserialized_args.request_event.data == HandoffRequest(target_agent="helper", reason="overflow")
assert deserialized_args.request_event.response_type is str
def test_process_request_info_event_passes_through_function_approval_request(self) -> None:
"""If the event data is already a function approval request, it is forwarded unchanged.
Tool-approval requests emitted by an inner agent surface as ``Content``
objects with ``user_input_request=True``. ``WorkflowAgent`` must not
re-wrap these inside a synthesized ``request_info`` function call;
instead it should return the original content as-is so callers can
respond with a matching ``function_approval_response``.
"""
executor = SimpleExecutor(id="executor1", response_text="Response")
workflow = WorkflowBuilder(start_executor=executor).build()
agent = WorkflowAgent(workflow=workflow, name="Approval Passthrough Agent")
approval_id = "approval-passthrough-1"
inner_function_call = Content.from_function_call(
call_id="tool-call-1",
name="delete_file",
arguments={"path": "/tmp/x"},
)
approval_request = Content.from_function_approval_request(
id=approval_id,
function_call=inner_function_call,
)
event = WorkflowEvent.request_info(
request_id=approval_id,
source_executor_id="executor1",
request_data=approval_request,
response_type=Content,
)
result = agent._process_request_info_event(event) # pyright: ignore[reportPrivateUsage]
# The original FunctionApprovalRequestContent is returned as-is — same
# instance, with the original tool name preserved (NOT replaced by the
# synthesized REQUEST_INFO_FUNCTION_NAME).
assert result is approval_request
assert result.type == "function_approval_request"
assert result.id == approval_id
assert result.user_input_request is True
assert result.function_call is inner_function_call # type: ignore[attr-defined]
assert result.function_call.name == "delete_file" # type: ignore[attr-defined]
assert result.function_call.name != WorkflowAgent.REQUEST_INFO_FUNCTION_NAME # type: ignore[attr-defined]
def test_extract_function_responses_passes_through_approval_response_approved(self) -> None:
"""A function_approval_response with approved=True is keyed by content.id and forwarded as-is.
After the refactor, ``WorkflowAgent`` no longer unwraps a synthesized
``request_info`` function call from approval responses — the response
content is routed straight back to the workflow under its own ``id``,
which matches the pending request id surfaced by
``_process_request_info_event``.
"""
executor = SimpleExecutor(id="executor1", response_text="Response")
workflow = WorkflowBuilder(start_executor=executor).build()
agent = WorkflowAgent(workflow=workflow, name="Approval Response Agent")
approval_id = "approval-response-approved-1"
inner_function_call = Content.from_function_call(
call_id="tool-call-1",
name="delete_file",
arguments={"path": "/tmp/x"},
)
approval_request = Content.from_function_approval_request(
id=approval_id,
function_call=inner_function_call,
)
approval_response = approval_request.to_function_approval_response(approved=True) # type: ignore[attr-defined]
message = Message(role="user", contents=[approval_response])
responses = agent._extract_function_responses([message]) # pyright: ignore[reportPrivateUsage]
assert set(responses.keys()) == {approval_id}
assert responses[approval_id] is approval_response
assert responses[approval_id].approved is True # type: ignore[attr-defined]
def test_extract_function_responses_passes_through_approval_response_denied(self) -> None:
"""A function_approval_response with approved=False is forwarded the same way as an approval.
Only the ``approved`` flag changes — routing back to the workflow is
identical for accept and reject paths.
"""
executor = SimpleExecutor(id="executor1", response_text="Response")
workflow = WorkflowBuilder(start_executor=executor).build()
agent = WorkflowAgent(workflow=workflow, name="Approval Response Agent")
approval_id = "approval-response-denied-1"
inner_function_call = Content.from_function_call(
call_id="tool-call-2",
name="send_email",
arguments={"to": "alice@example.com"},
)
approval_request = Content.from_function_approval_request(
id=approval_id,
function_call=inner_function_call,
)
approval_response = approval_request.to_function_approval_response(approved=False) # type: ignore[attr-defined]
message = Message(role="user", contents=[approval_response])
responses = agent._extract_function_responses([message]) # pyright: ignore[reportPrivateUsage]
assert set(responses.keys()) == {approval_id}
assert responses[approval_id] is approval_response
assert responses[approval_id].approved is False # type: ignore[attr-defined]
async def test_function_approval_request_flows_end_to_end_approved(self) -> None:
"""End-to-end: an executor emits a function_approval_request, the agent
forwards it unchanged, and an ``approved=True`` response resumes the workflow.
This exercises the full pass-through path:
``ctx.request_info(approval_content, ...)`` -> ``WorkflowAgent`` surfaces
the original ``FunctionApprovalRequestContent`` -> caller responds with a
``FunctionApprovalResponseContent`` -> ``WorkflowAgent`` routes it back
to the workflow which delivers it to the executor's ``@response_handler``.
"""
approval_id = "e2e-approval-1"
inner_function_call = Content.from_function_call(
call_id="tool-call-e2e-1",
name="delete_file",
arguments={"path": "/tmp/x"},
)
approval_request = Content.from_function_approval_request(
id=approval_id,
function_call=inner_function_call,
)
class ApprovalRequestingExecutor(Executor):
@handler
async def handle_message(self, _: list[Message], ctx: WorkflowContext) -> None:
await ctx.request_info(approval_request, Content, request_id=approval_id)
@response_handler
async def handle_response(
self,
original_request: Content,
response: Content,
ctx: WorkflowContext[Never, AgentResponse],
) -> None:
assert response.type == "function_approval_response"
assert response.id == approval_id # type: ignore[attr-defined]
approved = bool(response.approved) # type: ignore[attr-defined]
tool_name = original_request.function_call.name # type: ignore[attr-defined]
await ctx.yield_output(
AgentResponse(
messages=[
Message(
role="assistant",
contents=[Content.from_text(text=f"{tool_name} approved={approved}")],
)
]
)
)
executor = ApprovalRequestingExecutor(id="approval_requester")
workflow = WorkflowBuilder(start_executor=executor).build()
agent = WorkflowAgent(workflow=workflow, name="E2E Approval Agent")
# First run: workflow pauses with the approval request.
first = await agent.run("please delete it")
assert isinstance(first, AgentResponse)
forwarded = next(
(
c
for m in first.messages
for c in m.contents
if c.type == "function_approval_request" and c.id == approval_id
),
None,
)
assert forwarded is approval_request, "Approval request must surface unchanged"
pending = await workflow._runner_context.get_pending_request_info_events()
assert approval_id in pending
# Respond with approved=True.
approval_response = approval_request.to_function_approval_response(approved=True) # type: ignore[attr-defined]
final = await agent.run(Message(role="user", contents=[approval_response]))
assert isinstance(final, AgentResponse)
final_text = " ".join(m.text or "" for m in final.messages)
assert "delete_file approved=True" in final_text
pending = await workflow._runner_context.get_pending_request_info_events()
assert approval_id not in pending
async def test_function_approval_request_flows_end_to_end_denied(self) -> None:
"""End-to-end denied path: ``approved=False`` is delivered to the executor's
response handler so the workflow can branch on the rejection."""
approval_id = "e2e-approval-deny-1"
inner_function_call = Content.from_function_call(
call_id="tool-call-e2e-deny-1",
name="send_email",
arguments={"to": "alice@example.com"},
)
approval_request = Content.from_function_approval_request(
id=approval_id,
function_call=inner_function_call,
)
class ApprovalRequestingExecutor(Executor):
@handler
async def handle_message(self, _: list[Message], ctx: WorkflowContext) -> None:
await ctx.request_info(approval_request, Content, request_id=approval_id)
@response_handler
async def handle_response(
self,
original_request: Content,
response: Content,
ctx: WorkflowContext[Never, AgentResponse],
) -> None:
assert response.type == "function_approval_response"
assert response.id == approval_id # type: ignore[attr-defined]
approved = bool(response.approved) # type: ignore[attr-defined]
tool_name = original_request.function_call.name # type: ignore[attr-defined]
await ctx.yield_output(
AgentResponse(
messages=[
Message(
role="assistant",
contents=[Content.from_text(text=f"{tool_name} approved={approved}")],
)
]
)
)
executor = ApprovalRequestingExecutor(id="approval_requester_deny")
workflow = WorkflowBuilder(start_executor=executor).build()
agent = WorkflowAgent(workflow=workflow, name="E2E Approval Deny Agent")
first = await agent.run("please send")
assert isinstance(first, AgentResponse)
forwarded = next(
(
c
for m in first.messages
for c in m.contents
if c.type == "function_approval_request" and c.id == approval_id
),
None,
)
assert forwarded is approval_request
# Respond with approved=False.
approval_response = approval_request.to_function_approval_response(approved=False) # type: ignore[attr-defined]
final = await agent.run(Message(role="user", contents=[approval_response]))
assert isinstance(final, AgentResponse)
final_text = " ".join(m.text or "" for m in final.messages)
assert "send_email approved=False" in final_text
pending = await workflow._runner_context.get_pending_request_info_events()
assert approval_id not in pending
async def test_request_info_non_approval_flows_end_to_end(self) -> None:
"""End-to-end: when request data is not a function approval content, the
agent surfaces a synthesized ``function_call`` (name=REQUEST_INFO_FUNCTION_NAME)
and routes a matching ``function_result`` back to the executor.
"""
captured: dict[str, Any] = {}
class HandoffRequestingExecutor(Executor):
@handler
async def handle_message(self, _: list[Message], ctx: WorkflowContext) -> None:
await ctx.request_info(
HandoffRequest(target_agent="helper", reason="overflow"),
str,
)
@response_handler
async def handle_response(
self,
original_request: HandoffRequest,
response: str,
ctx: WorkflowContext[Never, AgentResponse],
) -> None:
captured["original"] = original_request
captured["response"] = response
await ctx.yield_output(
AgentResponse(
messages=[
Message(
role="assistant",
contents=[
Content.from_text(text=f"handoff to {original_request.target_agent}: {response}")
],
)
]
)
)
executor = HandoffRequestingExecutor(id="handoff_requester")
workflow = WorkflowBuilder(start_executor=executor).build()
agent = WorkflowAgent(workflow=workflow, name="E2E Handoff Agent")
# First run: workflow pauses with a synthesized request_info function_call.
first = await agent.run("start handoff")
assert isinstance(first, AgentResponse)
function_call = next(
(
c
for m in first.messages
for c in m.contents
if c.type == "function_call" and c.name == WorkflowAgent.REQUEST_INFO_FUNCTION_NAME
),
None,
)
assert function_call is not None, "Expected a synthesized request_info function_call"
assert function_call.call_id is not None
assert isinstance(function_call.arguments, dict)
request_id = function_call.arguments["request_id"]
assert function_call.call_id == request_id
request_payload = function_call.arguments["request_event"]
assert request_payload.get("type") == "request_info"
assert request_payload.get("data") == HandoffRequest(target_agent="helper", reason="overflow")
deserialized_args = WorkflowAgent.RequestInfoFunctionArgs.from_dict(function_call.arguments)
assert deserialized_args.request_id == request_id
assert isinstance(deserialized_args.request_event, WorkflowEvent)
assert deserialized_args.request_event.type == "request_info"
assert deserialized_args.request_event.data == HandoffRequest(target_agent="helper", reason="overflow")
assert deserialized_args.request_event.response_type is str
pending = await workflow._runner_context.get_pending_request_info_events()
assert request_id in pending
# Respond with a function_result keyed by the call_id.
function_result = Content.from_function_result(call_id=request_id, result="ok-do-it")
final = await agent.run(Message(role="user", contents=[function_result]))
assert isinstance(final, AgentResponse)
final_text = " ".join(m.text or "" for m in final.messages)
assert "handoff to helper: ok-do-it" in final_text
# The executor's response handler received the original request and the response.
assert isinstance(captured.get("original"), HandoffRequest)
assert captured["original"].target_agent == "helper"
assert captured["response"] == "ok-do-it"
pending = await workflow._runner_context.get_pending_request_info_events()
assert request_id not in pending
def test_workflow_as_agent_method(self) -> None:
"""Test that Workflow.as_agent() creates a properly configured WorkflowAgent."""
@@ -1592,3 +1946,406 @@ class TestWorkflowAgentMergeUpdates:
# Order: text (user), text (assistant), function_result (orphan at end)
assert content_types == ["text", "text", "function_result"]
class _ToolApprovalMockAgent(SupportsAgentRun):
"""Mock agent whose first run returns a FunctionApprovalRequestContent.
Subsequent runs (after receiving an approval response in the input messages)
return a final assistant text response that echoes the approved arguments.
This mirrors a real agent whose tool invocation requires user approval.
"""
def __init__(
self,
name: str,
*,
tool_name: str = "delete_file",
tool_arguments: dict[str, Any] | None = None,
approval_request_ids: Sequence[str] | None = None,
) -> None:
self.id = str(uuid.uuid4())
self.name = name
self.description: str | None = None
self._tool_name = tool_name
self._tool_arguments = tool_arguments or {"path": "/tmp/example"}
# Pre-allocated request ids so the test can verify what the WorkflowAgent forwards.
self._approval_request_ids: list[str] = list(approval_request_ids) if approval_request_ids else []
self.run_count = 0
# Inputs received on the most recent (continuation) run, for assertions.
self.last_run_messages: list[Message] = []
def create_session(self, **kwargs: Any) -> AgentSession:
return AgentSession()
def get_session(self, *, service_session_id: str, **kwargs: Any) -> AgentSession:
return AgentSession()
def _next_request_id(self) -> str:
if self._approval_request_ids:
return self._approval_request_ids.pop(0)
return str(uuid.uuid4())
def _build_approval_request(self) -> Content:
request_id = self._next_request_id()
function_call = Content.from_function_call(
call_id=request_id,
name=self._tool_name,
arguments=self._tool_arguments,
)
return Content.from_function_approval_request(id=request_id, function_call=function_call)
@overload
def run(
self,
messages: str | Content | Message | Sequence[str | Content | Message] | None = ...,
*,
stream: Literal[False] = ...,
session: AgentSession | None = ...,
**kwargs: Any,
) -> Awaitable[AgentResponse[Any]]: ...
@overload
def run(
self,
messages: str | Content | Message | Sequence[str | Content | Message] | None = ...,
*,
stream: Literal[True],
session: AgentSession | None = ...,
**kwargs: Any,
) -> ResponseStream[AgentResponseUpdate, AgentResponse[Any]]: ...
def run(
self,
messages: str | Content | Message | Sequence[str | Content | Message] | None = None,
*,
stream: bool = False,
session: AgentSession | None = None,
**kwargs: Any,
) -> Awaitable[AgentResponse] | ResponseStream[AgentResponseUpdate, AgentResponse]:
if stream:
return self._run_stream(messages=messages, session=session, **kwargs)
return self._run(messages=messages, session=session, **kwargs)
def _normalize(
self,
messages: str | Content | Message | Sequence[str | Content | Message] | None,
) -> list[Message]:
if messages is None:
return []
if isinstance(messages, str):
return [Message(role="user", contents=[Content.from_text(text=messages)])]
if isinstance(messages, Message):
return [messages]
if isinstance(messages, Content):
return [Message(role="user", contents=[messages])]
result: list[Message] = []
for item in messages:
if isinstance(item, Message):
result.append(item)
elif isinstance(item, Content):
result.append(Message(role="user", contents=[item]))
else:
result.append(Message(role="user", contents=[Content.from_text(text=item)]))
return result
def _approval_responses_in(self, messages: list[Message]) -> list[Content]:
approvals: list[Content] = []
for msg in messages:
for content in msg.contents:
if content.type == "function_approval_response":
approvals.append(content)
return approvals
async def _run(
self,
messages: str | Content | Message | Sequence[str | Content | Message] | None = None,
*,
session: AgentSession | None = None,
**kwargs: Any,
) -> AgentResponse:
normalized = self._normalize(messages)
self.last_run_messages = normalized
self.run_count += 1
approvals = self._approval_responses_in(normalized)
if approvals:
# Continuation: reflect approved arguments in the final response text.
approved_text = "; ".join(
f"approved={a.approved} id={a.id}" # type: ignore[attr-defined]
for a in approvals
)
return AgentResponse(messages=[Message("assistant", [Content.from_text(text=f"done ({approved_text})")])])
# First run: ask for tool approval.
approval = self._build_approval_request()
return AgentResponse(messages=[Message("assistant", [approval])])
def _run_stream(
self,
messages: str | Content | Message | Sequence[str | Content | Message] | None = None,
*,
session: AgentSession | None = None,
**kwargs: Any,
) -> ResponseStream[AgentResponseUpdate, AgentResponse]:
normalized = self._normalize(messages)
self.last_run_messages = normalized
self.run_count += 1
approvals = self._approval_responses_in(normalized)
async def _iter():
if approvals:
approved_text = "; ".join(
f"approved={a.approved} id={a.id}" # type: ignore[attr-defined]
for a in approvals
)
yield AgentResponseUpdate(
contents=[Content.from_text(text=f"done ({approved_text})")],
role="assistant",
author_name=self.name,
)
return
approval = self._build_approval_request()
yield AgentResponseUpdate(
contents=[approval],
role="assistant",
author_name=self.name,
)
return ResponseStream(_iter(), finalizer=AgentResponse.from_updates)
class TestWorkflowAgentToolApproval:
"""Tests for tool-approval requests bubbling through WorkflowAgent.
Covers the case where a workflow contains an AgentExecutor whose underlying
agent emits a FunctionApprovalRequestContent (tool needing user approval).
The WorkflowAgent must:
* forward the original FunctionApprovalRequestContent unchanged (no
wrapping inside a synthesized 'request_info' function call), and
* route a subsequent FunctionApprovalResponseContent back to the
AgentExecutor so the agent can resume.
"""
def _find_approval_request(
self,
contents: Sequence[Content],
tool_name: str,
) -> Content | None:
for content in contents:
if (
content.type == "function_approval_request"
and getattr(content.function_call, "name", None) == tool_name # type: ignore[attr-defined]
):
return content
return None
async def test_tool_approval_request_forwarded_unchanged(self) -> None:
"""The agent's FunctionApprovalRequestContent surfaces verbatim (not re-wrapped)."""
approval_id = "approval-abc-123"
mock_agent = _ToolApprovalMockAgent(
name="approval-agent",
tool_name="delete_file",
tool_arguments={"path": "/tmp/secret.txt"},
approval_request_ids=[approval_id],
)
@executor
async def start(messages: list[Message], ctx: WorkflowContext[AgentExecutorRequest]) -> None:
await ctx.send_message(AgentExecutorRequest(messages=messages, should_respond=True))
workflow = WorkflowBuilder(start_executor=start).add_edge(start, mock_agent).build()
agent = WorkflowAgent(workflow=workflow, name="Approval Test Agent")
result = await agent.run("please delete the file")
assert isinstance(result, AgentResponse)
# Locate the approval request emitted by the WorkflowAgent.
all_contents: list[Content] = [c for m in result.messages for c in m.contents]
approval = self._find_approval_request(all_contents, tool_name="delete_file")
assert approval is not None, "WorkflowAgent did not forward the tool approval request"
# The id and inner function_call must match what the underlying agent produced
# — i.e. the WorkflowAgent must NOT have re-wrapped it inside a synthesized
# 'request_info' approval request.
assert approval.id == approval_id
function_call = approval.function_call # type: ignore[attr-defined]
assert function_call is not None
assert function_call.name == "delete_file"
assert function_call.name != WorkflowAgent.REQUEST_INFO_FUNCTION_NAME
assert function_call.arguments == {"path": "/tmp/secret.txt"}
# The agent must be paused awaiting the approval response.
pending = await workflow._runner_context.get_pending_request_info_events()
assert approval_id in pending
async def test_tool_approval_request_forwarded_unchanged_streaming(self) -> None:
"""Streaming variant: the approval request is forwarded as-is in updates."""
approval_id = "approval-stream-1"
mock_agent = _ToolApprovalMockAgent(
name="approval-agent-stream",
tool_name="send_email",
tool_arguments={"to": "alice@example.com"},
approval_request_ids=[approval_id],
)
@executor
async def start(messages: list[Message], ctx: WorkflowContext[AgentExecutorRequest]) -> None:
await ctx.send_message(AgentExecutorRequest(messages=messages, should_respond=True))
workflow = WorkflowBuilder(start_executor=start).add_edge(start, mock_agent).build()
agent = WorkflowAgent(workflow=workflow, name="Approval Stream Agent")
updates: list[AgentResponseUpdate] = []
async for update in agent.run("hi", stream=True):
updates.append(update)
approval_updates = [u for u in updates if any(c.type == "function_approval_request" for c in u.contents)]
assert approval_updates, "Streaming did not surface a tool approval request"
approval = self._find_approval_request(approval_updates[-1].contents, tool_name="send_email")
assert approval is not None
assert approval.id == approval_id
function_call = approval.function_call # type: ignore[attr-defined]
assert function_call is not None
assert function_call.name == "send_email"
assert function_call.name != WorkflowAgent.REQUEST_INFO_FUNCTION_NAME
assert function_call.arguments == {"to": "alice@example.com"}
async def test_tool_approval_response_resumes_agent(self) -> None:
"""Sending the approval response back resumes the agent and clears pending requests."""
approval_id = "approval-resume-1"
mock_agent = _ToolApprovalMockAgent(
name="approval-resume-agent",
tool_name="delete_file",
tool_arguments={"path": "/tmp/x"},
approval_request_ids=[approval_id],
)
@executor
async def start(messages: list[Message], ctx: WorkflowContext[AgentExecutorRequest]) -> None:
await ctx.send_message(AgentExecutorRequest(messages=messages, should_respond=True))
workflow = WorkflowBuilder(start_executor=start).add_edge(start, mock_agent).build()
agent = WorkflowAgent(workflow=workflow, name="Approval Resume Agent")
first_result = await agent.run("delete it")
approval = self._find_approval_request(
[c for m in first_result.messages for c in m.contents],
tool_name="delete_file",
)
assert approval is not None
assert mock_agent.run_count == 1
# Build the approval response. NOTE: the inner function_call's name is the
# original tool name ('delete_file'), NOT 'request_info'. This exercises the
# branch in WorkflowAgent._extract_function_responses that routes raw
# tool-approval responses straight through using content.id.
approval_response = approval.to_function_approval_response(approved=True) # type: ignore[attr-defined]
response_message = Message(role="user", contents=[approval_response])
final_result = await agent.run(response_message)
assert isinstance(final_result, AgentResponse)
# The mock agent should have been invoked a second time and seen the
# approval response in its inputs.
assert mock_agent.run_count == 2
approvals_seen = [
c for m in mock_agent.last_run_messages for c in m.contents if c.type == "function_approval_response"
]
assert len(approvals_seen) == 1
assert approvals_seen[0].id == approval_id # type: ignore[attr-defined]
assert approvals_seen[0].approved is True # type: ignore[attr-defined]
# The pending approval should now be cleared.
pending = await workflow._runner_context.get_pending_request_info_events()
assert approval_id not in pending
# The final assistant message reflects the resumption.
final_text = " ".join(m.text or "" for m in final_result.messages)
assert "done" in final_text
assert approval_id in final_text
async def test_tool_approval_response_rejected_resumes_agent(self) -> None:
"""Rejection path: ``approved=False`` is forwarded to the inner agent and clears the pending request.
The WorkflowAgent must route a rejection response back to the paused
``AgentExecutor`` exactly the same way as an approval — only the
``approved`` flag differs. The inner agent decides what to do with it.
"""
approval_id = "approval-reject-1"
mock_agent = _ToolApprovalMockAgent(
name="approval-reject-agent",
tool_name="delete_file",
tool_arguments={"path": "/tmp/x"},
approval_request_ids=[approval_id],
)
@executor
async def start(messages: list[Message], ctx: WorkflowContext[AgentExecutorRequest]) -> None:
await ctx.send_message(AgentExecutorRequest(messages=messages, should_respond=True))
workflow = WorkflowBuilder(start_executor=start).add_edge(start, mock_agent).build()
agent = WorkflowAgent(workflow=workflow, name="Approval Reject Agent")
first_result = await agent.run("delete it")
approval = self._find_approval_request(
[c for m in first_result.messages for c in m.contents],
tool_name="delete_file",
)
assert approval is not None
assert mock_agent.run_count == 1
# Reject the tool invocation.
approval_response = approval.to_function_approval_response(approved=False) # type: ignore[attr-defined]
response_message = Message(role="user", contents=[approval_response])
final_result = await agent.run(response_message)
assert isinstance(final_result, AgentResponse)
# The inner agent must have been resumed and seen ``approved=False``.
assert mock_agent.run_count == 2
approvals_seen = [
c for m in mock_agent.last_run_messages for c in m.contents if c.type == "function_approval_response"
]
assert len(approvals_seen) == 1
assert approvals_seen[0].id == approval_id # type: ignore[attr-defined]
assert approvals_seen[0].approved is False # type: ignore[attr-defined]
# Pending approval cleared regardless of approve/reject.
pending = await workflow._runner_context.get_pending_request_info_events()
assert approval_id not in pending
# The final assistant message reflects the rejection.
final_text = " ".join(m.text or "" for m in final_result.messages)
assert "approved=False" in final_text
assert approval_id in final_text
async def test_tool_approval_request_id_matches_pending_request(self) -> None:
"""The approval request id surfaced by WorkflowAgent matches the workflow's pending request id.
This guards the AgentExecutor change that forwards
request_id=user_input_request.id to ctx.request_info(...), which is what
allows the response routed back via WorkflowAgent to resolve the pending
request without an id-mismatch error.
"""
approval_id = "approval-id-match-1"
mock_agent = _ToolApprovalMockAgent(
name="approval-id-match-agent",
approval_request_ids=[approval_id],
)
@executor
async def start(messages: list[Message], ctx: WorkflowContext[AgentExecutorRequest]) -> None:
await ctx.send_message(AgentExecutorRequest(messages=messages, should_respond=True))
workflow = WorkflowBuilder(start_executor=start).add_edge(start, mock_agent).build()
agent = WorkflowAgent(workflow=workflow, name="Approval Id Agent")
await agent.run("go")
pending = await workflow._runner_context.get_pending_request_info_events()
# The agent's approval id is used as the workflow's pending request id.
assert list(pending.keys()) == [approval_id]
@@ -0,0 +1,149 @@
# Copyright (c) Microsoft. All rights reserved.
"""Tests for the ``Workflow.status`` property."""
from dataclasses import dataclass
import pytest
from agent_framework import (
Executor,
Workflow,
WorkflowBuilder,
WorkflowContext,
WorkflowEvent,
WorkflowRunState,
handler,
response_handler,
)
from agent_framework._workflows._executor import Executor as _Executor
from agent_framework._workflows._request_info_mixin import RequestInfoMixin
class PassThroughExecutor(Executor):
"""Executor that yields its input as a workflow output and stops."""
@handler
async def passthrough(self, msg: str, ctx: WorkflowContext[str, str]) -> None:
await ctx.yield_output(msg)
class FailingExecutor(Executor):
"""Executor that raises at runtime to drive the FAILED status."""
@handler
async def fail(self, msg: int, ctx: WorkflowContext) -> None: # pragma: no cover - invoked via workflow
raise RuntimeError("boom")
@dataclass
class _ApprovalRequest:
prompt: str
request_id: str = ""
def __post_init__(self) -> None:
if not self.request_id:
import uuid
self.request_id = str(uuid.uuid4())
class ApprovalExecutor(_Executor, RequestInfoMixin):
"""Executor that issues a single request_info call and finalizes on response."""
def __init__(self, id: str = "approval"):
super().__init__(id=id)
@handler
async def start(self, message: str, ctx: WorkflowContext[str, str]) -> None:
await ctx.request_info(_ApprovalRequest(prompt=message), bool)
@response_handler
async def on_response(
self, original_request: _ApprovalRequest, approved: bool, ctx: WorkflowContext[str, str]
) -> None:
await ctx.yield_output(f"approved={approved}")
def _build_passthrough_workflow() -> Workflow:
executor = PassThroughExecutor(id="p")
return WorkflowBuilder(start_executor=executor, output_from=[executor]).build()
def _build_failing_workflow() -> Workflow:
# FailingExecutor has no workflow_output_types, so we leave designation
# implicit; the deprecation warning is filtered at call sites that need it.
return WorkflowBuilder(start_executor=FailingExecutor(id="f")).build()
def _build_approval_workflow() -> Workflow:
executor = ApprovalExecutor(id="approval")
return WorkflowBuilder(start_executor=executor, output_from=[executor]).build()
async def test_status_default_is_idle_before_first_run():
wf = _build_passthrough_workflow()
assert wf.status is WorkflowRunState.IDLE
async def test_status_is_idle_after_successful_run():
wf = _build_passthrough_workflow()
await wf.run("hello")
assert wf.status is WorkflowRunState.IDLE
async def test_status_is_failed_after_failure():
wf = _build_failing_workflow()
with pytest.raises(RuntimeError, match="boom"):
await wf.run(0)
assert wf.status is WorkflowRunState.FAILED
async def test_status_transitions_during_streaming_run():
"""Workflow.status mirrors the most recent emitted status event."""
wf = _build_passthrough_workflow()
observed: list[WorkflowRunState] = []
async for event in wf.run("hi", stream=True):
if isinstance(event, WorkflowEvent) and event.type == "status":
# By the time a status event surfaces to the consumer, the property
# must already reflect that state (updated in lockstep with emission).
assert wf.status == event.state
observed.append(event.state) # type: ignore
# IN_PROGRESS must precede IDLE; both must appear.
assert WorkflowRunState.IN_PROGRESS in observed
assert observed[-1] is WorkflowRunState.IDLE
assert wf.status is WorkflowRunState.IDLE
async def test_status_idle_with_pending_requests_then_resolves_to_idle():
wf = _build_approval_workflow()
request_event: WorkflowEvent | None = None
async for event in wf.run("please approve", stream=True):
if isinstance(event, WorkflowEvent) and event.type == "request_info":
request_event = event
assert request_event is not None
assert wf.status is WorkflowRunState.IDLE_WITH_PENDING_REQUESTS
async for _ in wf.run(stream=True, responses={request_event.request_id: True}):
pass
assert wf.status is WorkflowRunState.IDLE
async def test_status_in_progress_pending_requests_observed_mid_run():
"""While streaming, status reaches IN_PROGRESS_PENDING_REQUESTS after a request_info event."""
wf = _build_approval_workflow()
seen_states: list[WorkflowRunState] = []
async for event in wf.run("please approve", stream=True):
if isinstance(event, WorkflowEvent) and event.type == "status":
seen_states.append(event.state) # type: ignore
assert WorkflowRunState.IN_PROGRESS in seen_states
assert WorkflowRunState.IN_PROGRESS_PENDING_REQUESTS in seen_states
assert seen_states[-1] is WorkflowRunState.IDLE_WITH_PENDING_REQUESTS
assert wf.status is WorkflowRunState.IDLE_WITH_PENDING_REQUESTS
+1 -1
View File
@@ -26,7 +26,7 @@ dependencies = [
"agent-framework-core>=1.8.0,<2",
"agent-framework-openai>=1.8.0,<2",
"azure-ai-inference>=1.0.0b9,<1.0.0b10",
"azure-ai-projects>=2.1.0,<3.0",
"azure-ai-projects>=2.2.0,<3.0",
]
[tool.uv]
@@ -567,7 +567,7 @@ class ResponsesHostServer(ResponsesAgentServerHost):
by the hosting infrastructure or files will be preserved upon deactivation.
"""
input_items = await context.get_input_items()
input_messages = await _items_to_messages(input_items)
input_messages = await _items_to_messages(input_items, approval_storage=self._approval_storage)
is_streaming_request = request.stream is not None and request.stream is True
_, are_options_set = _to_chat_options(request)
@@ -664,7 +664,11 @@ class ResponsesHostServer(ResponsesAgentServerHost):
checkpoint_storage=write_storage,
)
async for item in _to_outputs_for_messages(response_event_stream, response.messages):
async for item in _to_outputs_for_messages(
response_event_stream,
response.messages,
approval_storage=self._approval_storage,
):
yield item
await self._delete_not_latest_checkpoints(write_storage, self._agent.workflow.name)
@@ -685,7 +689,9 @@ class ResponsesHostServer(ResponsesAgentServerHost):
for event in tracker.handle(content):
yield event
if tracker.needs_async:
async for item in _to_outputs(response_event_stream, content):
async for item in _to_outputs(
response_event_stream, content, approval_storage=self._approval_storage
):
yield item
tracker.needs_async = False
@@ -11,24 +11,33 @@ the registered _handle_create handler.
from __future__ import annotations
import json
from collections.abc import AsyncIterator, Callable
import uuid
from collections.abc import AsyncIterator, Awaitable, Callable, Sequence
from dataclasses import dataclass
from typing import Literal, overload
from unittest.mock import AsyncMock, MagicMock, patch
import httpx
import pytest
from agent_framework import (
AgentExecutorRequest,
AgentResponse,
AgentResponseUpdate,
AgentSession,
Content,
FileCheckpointStorage,
HistoryProvider,
Message,
RawAgent,
ResponseStream,
SupportsAgentRun,
WorkflowAgent,
WorkflowBuilder,
WorkflowCheckpoint,
WorkflowCheckpointException,
WorkflowContext,
WorkflowMessage,
executor,
)
from azure.ai.agentserver.responses import InMemoryResponseProvider
from mcp import McpError
@@ -102,7 +111,7 @@ def _make_agent(
return agent
def _make_server(agent: MagicMock, **kwargs: Any) -> ResponsesHostServer:
def _make_server(agent: Any, **kwargs: Any) -> ResponsesHostServer:
"""Create a ResponsesHostServer with an in-memory store."""
return ResponsesHostServer(agent, store=InMemoryResponseProvider(), **kwargs)
@@ -3469,3 +3478,498 @@ class TestOAuthConsentSurfacing:
# endregion
# region Workflow agent hosting (end-to-end)
class _ToolApprovalWorkflowAgentMock(SupportsAgentRun):
"""Inner agent for a hosted ``WorkflowAgent`` whose first run emits a
``FunctionApprovalRequestContent`` and whose follow-up run (after
receiving a ``FunctionApprovalResponseContent`` in its inputs) returns a
final assistant text response.
Mirrors a real agent whose tool invocation requires user approval. Used
here to exercise the full HTTP pipeline through ``ResponsesHostServer``
when the hosted agent is a ``WorkflowAgent`` containing a tool-approval
flow.
"""
def __init__(
self,
name: str,
*,
tool_name: str = "delete_file",
tool_arguments: dict[str, Any] | None = None,
approval_request_ids: Sequence[str] | None = None,
final_text: str = "done",
) -> None:
self.id = str(uuid.uuid4())
self.name = name
self.description: str | None = None
self._tool_name = tool_name
self._tool_arguments = tool_arguments or {"path": "/tmp/example"}
self._approval_request_ids: list[str] = list(approval_request_ids) if approval_request_ids else []
self._final_text = final_text
self.run_count = 0
self.last_run_messages: list[Message] = []
def create_session(self, **kwargs: Any) -> AgentSession:
return AgentSession()
def get_session(self, *, service_session_id: str, **kwargs: Any) -> AgentSession:
return AgentSession()
def _next_request_id(self) -> str:
# Stable across calls: when the workflow checkpoint round-trips through
# restore, ``AgentExecutor`` re-invokes the inner agent during replay.
# We must surface the *same* approval request id on each invocation so
# the workflow's pending-request id matches the id the test echoes
# back as ``mcp_approval_response``.
if self._approval_request_ids:
return self._approval_request_ids[0]
return str(uuid.uuid4())
def _build_approval_request(self) -> Content:
request_id = self._next_request_id()
function_call = Content.from_function_call(
call_id=request_id,
name=self._tool_name,
arguments=self._tool_arguments,
additional_properties={"server_label": "test_server"},
)
return Content.from_function_approval_request(id=request_id, function_call=function_call)
@overload
def run(
self,
messages: str | Content | Message | Sequence[str | Content | Message] | None = ...,
*,
stream: Literal[False] = ...,
session: AgentSession | None = ...,
**kwargs: Any,
) -> Awaitable[AgentResponse[Any]]: ...
@overload
def run(
self,
messages: str | Content | Message | Sequence[str | Content | Message] | None = ...,
*,
stream: Literal[True],
session: AgentSession | None = ...,
**kwargs: Any,
) -> ResponseStream[AgentResponseUpdate, AgentResponse[Any]]: ...
def run(
self,
messages: str | Content | Message | Sequence[str | Content | Message] | None = None,
*,
stream: bool = False,
session: AgentSession | None = None,
**kwargs: Any,
) -> Awaitable[AgentResponse] | ResponseStream[AgentResponseUpdate, AgentResponse]:
if stream:
return self._run_stream(messages=messages, **kwargs)
return self._run(messages=messages, **kwargs)
@staticmethod
def _normalize(
messages: str | Content | Message | Sequence[str | Content | Message] | None,
) -> list[Message]:
if messages is None:
return []
if isinstance(messages, str):
return [Message(role="user", contents=[Content.from_text(text=messages)])]
if isinstance(messages, Message):
return [messages]
if isinstance(messages, Content):
return [Message(role="user", contents=[messages])]
result: list[Message] = []
for item in messages:
if isinstance(item, Message):
result.append(item)
elif isinstance(item, Content):
result.append(Message(role="user", contents=[item]))
else:
result.append(Message(role="user", contents=[Content.from_text(text=item)]))
return result
@staticmethod
def _approval_responses_in(messages: list[Message]) -> list[Content]:
return [c for m in messages for c in m.contents if c.type == "function_approval_response"]
async def _run(
self,
messages: str | Content | Message | Sequence[str | Content | Message] | None = None,
**kwargs: Any,
) -> AgentResponse:
normalized = self._normalize(messages)
self.last_run_messages = normalized
self.run_count += 1
if self._approval_responses_in(normalized):
return AgentResponse(messages=[Message("assistant", [Content.from_text(text=self._final_text)])])
approval = self._build_approval_request()
return AgentResponse(messages=[Message("assistant", [approval])])
def _run_stream(
self,
messages: str | Content | Message | Sequence[str | Content | Message] | None = None,
**kwargs: Any,
) -> ResponseStream[AgentResponseUpdate, AgentResponse]:
normalized = self._normalize(messages)
self.last_run_messages = normalized
self.run_count += 1
approvals = self._approval_responses_in(normalized)
async def _iter() -> AsyncIterator[AgentResponseUpdate]:
if approvals:
yield AgentResponseUpdate(
contents=[Content.from_text(text=self._final_text)],
role="assistant",
author_name=self.name,
)
return
yield AgentResponseUpdate(
contents=[self._build_approval_request()],
role="assistant",
author_name=self.name,
)
return ResponseStream(_iter(), finalizer=AgentResponse.from_updates)
def _build_text_workflow_agent(text: str) -> WorkflowAgent:
"""Build a minimal ``WorkflowAgent`` whose inner agent emits a fixed text."""
class _TextAgent(SupportsAgentRun):
def __init__(self, name: str, text: str) -> None:
self.id = str(uuid.uuid4())
self.name = name
self.description: str | None = None
self._text = text
def create_session(self, **kwargs: Any) -> AgentSession:
return AgentSession()
def get_session(self, *, service_session_id: str, **kwargs: Any) -> AgentSession:
return AgentSession()
@overload
def run(
self,
messages: Any = ...,
*,
stream: Literal[False] = ...,
session: AgentSession | None = ...,
**kwargs: Any,
) -> Awaitable[AgentResponse[Any]]: ...
@overload
def run(
self,
messages: Any = ...,
*,
stream: Literal[True],
session: AgentSession | None = ...,
**kwargs: Any,
) -> ResponseStream[AgentResponseUpdate, AgentResponse[Any]]: ...
def run(
self,
messages: Any = None,
*,
stream: bool = False,
session: AgentSession | None = None,
**kwargs: Any,
) -> Awaitable[AgentResponse] | ResponseStream[AgentResponseUpdate, AgentResponse]:
text = self._text
name = self.name
async def _aresult() -> AgentResponse:
return AgentResponse(messages=[Message("assistant", [Content.from_text(text=text)])])
async def _aiter() -> AsyncIterator[AgentResponseUpdate]:
yield AgentResponseUpdate(
contents=[Content.from_text(text=text)],
role="assistant",
author_name=name,
)
if stream:
return ResponseStream(_aiter(), finalizer=AgentResponse.from_updates)
return _aresult()
inner = _TextAgent("text-agent", text)
@executor
async def start(messages: list[Message], ctx: WorkflowContext[AgentExecutorRequest]) -> None:
await ctx.send_message(AgentExecutorRequest(messages=messages, should_respond=True))
workflow = WorkflowBuilder(start_executor=start).add_edge(start, inner).build()
return WorkflowAgent(workflow=workflow, name="Text Workflow Agent")
def _build_approval_workflow_agent(
*,
approval_request_id: str,
tool_name: str = "delete_file",
tool_arguments: dict[str, Any] | None = None,
final_text: str = "done",
) -> tuple[WorkflowAgent, _ToolApprovalWorkflowAgentMock]:
"""Build a ``WorkflowAgent`` whose inner agent emits a tool approval request."""
mock_agent = _ToolApprovalWorkflowAgentMock(
name="approval-agent",
tool_name=tool_name,
tool_arguments=tool_arguments or {"path": "/tmp/secret.txt"},
approval_request_ids=[approval_request_id],
final_text=final_text,
)
@executor
async def start(messages: list[Message], ctx: WorkflowContext[AgentExecutorRequest]) -> None:
await ctx.send_message(AgentExecutorRequest(messages=messages, should_respond=True))
workflow = WorkflowBuilder(start_executor=start).add_edge(start, mock_agent).build()
workflow_agent = WorkflowAgent(workflow=workflow, name="Approval Workflow Agent")
return workflow_agent, mock_agent
class TestWorkflowAgentHosting:
"""End-to-end HTTP tests for ``ResponsesHostServer`` hosting a ``WorkflowAgent``.
These tests drive ``_handle_inner_workflow`` through the ASGI stack:
they exercise checkpoint write/restore (multi-turn) and the
tool-approval round-trip path, which is the primary differentiator
relative to the regular agent path.
"""
async def test_basic_text_response(self) -> None:
workflow_agent = _build_text_workflow_agent("hello from workflow")
server = _make_server(workflow_agent)
resp = await _post(server, input_text="hi", stream=False)
assert resp.status_code == 200
body = resp.json()
assert body["status"] == "completed"
text_found = any(
part.get("type") == "output_text" and part.get("text") == "hello from workflow"
for item in body["output"]
if item["type"] == "message"
for part in item.get("content", [])
)
assert text_found, f"Expected workflow output text in {body['output']}"
async def test_basic_text_response_streaming(self) -> None:
workflow_agent = _build_text_workflow_agent("hello stream")
server = _make_server(workflow_agent)
resp = await _post(server, input_text="hi", stream=True)
assert resp.status_code == 200
events = _parse_sse_events(resp.text)
types = _sse_event_types(events)
assert types[0] == "response.created"
assert types[-1] == "response.completed"
assert "response.output_text.delta" in types
text_done = [e for e in events if e["event"] == "response.output_text.done"]
assert any(e["data"]["text"] == "hello stream" for e in text_done)
async def test_non_streaming_emits_mcp_approval_request_and_persists_to_storage(self) -> None:
workflow_agent, mock_agent = _build_approval_workflow_agent(approval_request_id="apr_wf_ns")
server = _make_server(workflow_agent)
resp = await _post(server, stream=False)
assert resp.status_code == 200
body = resp.json()
assert body["status"] == "completed"
approval_items = [it for it in body["output"] if it["type"] == "mcp_approval_request"]
assert len(approval_items) == 1
assert approval_items[0]["name"] == "delete_file"
assert approval_items[0]["server_label"] == "test_server"
approval_request_id = approval_items[0]["id"]
# The id surfaced over the wire is generated by the response stream
# builder; the original approval ``Content`` (carrying the inner
# ``function_call``) must be persisted under that id so the next
# turn can reconstruct it.
loaded = await server._approval_storage.load_approval_request( # pyright: ignore[reportPrivateUsage]
approval_request_id
)
assert loaded.type == "function_approval_request"
assert loaded.function_call.name == "delete_file" # type: ignore[attr-defined]
assert mock_agent.run_count == 1
async def test_streaming_emits_mcp_approval_request_and_persists_to_storage(self) -> None:
workflow_agent, mock_agent = _build_approval_workflow_agent(approval_request_id="apr_wf_st")
server = _make_server(workflow_agent)
resp = await _post(server, stream=True)
assert resp.status_code == 200
events = _parse_sse_events(resp.text)
types = _sse_event_types(events)
assert types[0] == "response.created"
assert types[-1] == "response.completed"
approval_request_id: str | None = None
for e in events:
if e["event"] != "response.output_item.added":
continue
item = e["data"].get("item") or {}
if item.get("type") == "mcp_approval_request":
approval_request_id = item.get("id")
break
assert approval_request_id is not None
loaded = await server._approval_storage.load_approval_request( # pyright: ignore[reportPrivateUsage]
approval_request_id
)
assert loaded.type == "function_approval_request"
assert mock_agent.run_count == 1
async def test_round_trip_approval_response_resumes_workflow_agent(self) -> None:
"""Two-turn HTTP round-trip:
Turn 1 emits ``mcp_approval_request`` and writes a workflow
checkpoint under the response id. Turn 2 sends the
``mcp_approval_response`` with ``previous_response_id`` set, so the
host restores the checkpoint, the WorkflowAgent routes the
approval response back to the paused inner agent, and the inner
agent emits the final assistant text.
"""
workflow_agent, mock_agent = _build_approval_workflow_agent(
approval_request_id="apr_wf_rt",
final_text="done with approval",
)
server = _make_server(workflow_agent)
first = await _post(server, stream=False)
assert first.status_code == 200
first_body = first.json()
first_response_id = first_body["id"]
approval_items = [it for it in first_body["output"] if it["type"] == "mcp_approval_request"]
assert len(approval_items) == 1
approval_request_id = approval_items[0]["id"]
assert mock_agent.run_count == 1
second_payload: dict[str, Any] = {
"model": "test-model",
"input": [
{
"type": "mcp_approval_response",
"approval_request_id": approval_request_id,
"approve": True,
}
],
"stream": False,
"previous_response_id": first_response_id,
}
second = await _post_json(server, second_payload)
assert second.status_code == 200
second_body = second.json()
assert second_body["status"] == "completed"
# The inner agent must have been resumed (restore replay + new turn).
# Restore call is a no-op for the mock (no input); the new-turn call
# delivers the approval response, so run_count grows by at least 1.
assert mock_agent.run_count >= 2
# The final assistant text from the resumed inner agent surfaces in
# the HTTP output.
text_pieces = [
part.get("text", "")
for item in second_body["output"]
if item["type"] == "message"
for part in item.get("content", [])
if part.get("type") == "output_text"
]
assert any("done with approval" in t for t in text_pieces), (
f"expected resumed workflow output, got {second_body['output']}"
)
# The new-turn invocation of the inner agent must have received the
# approval response routed back through WorkflowAgent.
approval_responses = [
c for m in mock_agent.last_run_messages for c in m.contents if c.type == "function_approval_response"
]
assert len(approval_responses) == 1
assert approval_responses[0].approved is True # type: ignore[attr-defined]
async def test_round_trip_approval_response_streaming(self) -> None:
"""Streaming variant of the round-trip: turn 2 is requested with
``stream=true`` and surfaces the resumed text as SSE events."""
workflow_agent, mock_agent = _build_approval_workflow_agent(
approval_request_id="apr_wf_rt_st",
final_text="streamed-done",
)
server = _make_server(workflow_agent)
first = await _post(server, stream=False)
first_body = first.json()
first_response_id = first_body["id"]
approval_request_id = next(it["id"] for it in first_body["output"] if it["type"] == "mcp_approval_request")
second = await _post_json(
server,
{
"model": "test-model",
"input": [
{
"type": "mcp_approval_response",
"approval_request_id": approval_request_id,
"approve": True,
}
],
"stream": True,
"previous_response_id": first_response_id,
},
)
assert second.status_code == 200
events = _parse_sse_events(second.text)
types = _sse_event_types(events)
assert types[0] == "response.created"
assert types[-1] == "response.completed"
text_done = [e for e in events if e["event"] == "response.output_text.done"]
assert any("streamed-done" in e["data"]["text"] for e in text_done)
assert mock_agent.run_count >= 2
async def test_round_trip_approval_response_rejected(self) -> None:
"""Sending ``approve=False`` must surface as ``approved=False`` to the
inner agent on resume."""
workflow_agent, mock_agent = _build_approval_workflow_agent(
approval_request_id="apr_wf_reject",
final_text="acknowledged",
)
server = _make_server(workflow_agent)
first = await _post(server, stream=False)
first_body = first.json()
first_response_id = first_body["id"]
approval_request_id = next(it["id"] for it in first_body["output"] if it["type"] == "mcp_approval_request")
second = await _post_json(
server,
{
"model": "test-model",
"input": [
{
"type": "mcp_approval_response",
"approval_request_id": approval_request_id,
"approve": False,
}
],
"stream": False,
"previous_response_id": first_response_id,
},
)
assert second.status_code == 200
approval_responses = [
c for m in mock_agent.last_run_messages for c in m.contents if c.type == "function_approval_response"
]
assert len(approval_responses) == 1
assert approval_responses[0].approved is False # type: ignore[attr-defined]
# endregion
+4
View File
@@ -10,6 +10,10 @@ pip install agent-framework-gemini --pre
The Gemini integration enables Microsoft Agent Framework applications to call Google Gemini models with familiar chat abstractions, including streaming, tool/function calling, and structured output.
## Structured Output
Gemini structured output can be configured with either a Pydantic model in `response_format`, a JSON schema mapping in `response_format`, or a Gemini-specific `response_schema`. Declarative agents that define `outputSchema` pass that schema through `response_format`.
## Authentication
The connector supports both `google-genai` authentication modes.
@@ -109,8 +109,8 @@ class GeminiChatOptions(ChatOptions[ResponseModelT], Generic[ResponseModelT], to
or ``types.Tool`` objects returned by ``get_code_interpreter_tool``, ``get_web_search_tool``,
``get_mcp_tool``, ``get_file_search_tool``, or ``get_maps_grounding_tool``.
tool_choice: How the model picks a tool. One of ``'auto'``, ``'none'``, or ``'required'``.
response_format: Pydantic model type for structured JSON output. The response text is
parsed into the model and exposed via ``ChatResponse.value``.
response_format: Pydantic model type or JSON schema mapping for structured JSON output.
The response text is parsed and exposed via ``ChatResponse.value``.
instructions: Extra system-level instructions prepended to the system message.
Not supported, and passing these raises a type error:
@@ -255,6 +255,29 @@ _OPTION_CONSUMED_KEYS: frozenset[str] = frozenset({
_OPTION_EXCLUDE_KEYS: frozenset[str] = _OPTION_EXPLICIT_KEYS | _OPTION_CONSUMED_KEYS
_JSON_SCHEMA_TYPES: frozenset[str] = frozenset({
"array",
"boolean",
"integer",
"null",
"number",
"object",
"string",
})
_JSON_SCHEMA_KEYWORDS: frozenset[str] = frozenset({
"$defs",
"additionalProperties",
"allOf",
"anyOf",
"enum",
"items",
"oneOf",
"properties",
"required",
"type",
})
_FINISH_REASON_MAP: dict[str, FinishReasonLiteral] = {
"STOP": "stop",
"MAX_TOKENS": "length",
@@ -747,9 +770,13 @@ class RawGeminiChatClient(
continue
kwargs[_OPTION_TRANSLATIONS.get(key, key)] = value
if options.get("response_format") or options.get("response_schema"):
response_format = options.get("response_format")
response_schema = options.get("response_schema")
if response_format is not None or response_schema is not None:
kwargs["response_mime_type"] = "application/json"
if schema := options.get("response_schema"):
if response_schema is not None:
kwargs["response_schema"] = response_schema
elif (schema := self._extract_response_schema(response_format)) is not None:
kwargs["response_schema"] = schema
if tools := self._prepare_tools(options):
kwargs["tools"] = tools
@@ -762,6 +789,48 @@ class RawGeminiChatClient(
return types.GenerateContentConfig(**kwargs)
@staticmethod
def _extract_response_schema(response_format: Any) -> dict[str, Any] | None:
"""Extract a Gemini response schema from supported mapping response_format shapes."""
if not isinstance(response_format, Mapping):
return None
mapping = cast("Mapping[str, Any]", response_format)
if (nested := RawGeminiChatClient._extract_response_schema(mapping.get("format"))) is not None:
return nested
json_schema = mapping.get("json_schema")
if isinstance(json_schema, Mapping):
schema = cast("Mapping[str, Any]", json_schema).get("schema")
if isinstance(schema, Mapping):
return dict(cast("Mapping[str, Any]", schema))
schema = mapping.get("schema")
if isinstance(schema, Mapping):
return dict(cast("Mapping[str, Any]", schema))
if RawGeminiChatClient._is_json_schema_mapping(mapping):
return dict(mapping)
return None
@staticmethod
def _is_json_schema_mapping(value: Mapping[str, Any]) -> bool:
"""Return True when a mapping appears to be a JSON Schema rather than a response-format envelope."""
if not any(keyword in value for keyword in _JSON_SCHEMA_KEYWORDS):
return False
schema_type = value.get("type")
if schema_type is None:
return True
if isinstance(schema_type, str):
return schema_type in _JSON_SCHEMA_TYPES
if isinstance(schema_type, Sequence) and not isinstance(schema_type, (str, bytes)):
entries = cast("Sequence[object]", schema_type)
return all(isinstance(item, str) and item in _JSON_SCHEMA_TYPES for item in entries)
return False
def _prepare_tools(self, options: Mapping[str, Any]) -> list[types.Tool] | None:
"""Translate the framework tool list into Gemini API tool objects.
@@ -9,7 +9,7 @@ from typing import Any
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from agent_framework import Content, FunctionTool, Message
from agent_framework import Agent, Content, FunctionTool, Message
from google.genai import types
from pydantic import BaseModel
@@ -915,6 +915,20 @@ async def test_response_format_populates_value_on_chat_response() -> None:
assert response.value == Reply(text="hello")
async def test_response_format_mapping_populates_value_on_chat_response() -> None:
"""When response_format is a JSON schema mapping, ChatResponse.value must parse the response text."""
client, mock = _make_gemini_client()
mock.aio.models.generate_content = AsyncMock(return_value=_make_response([_make_part(text='{"text": "hello"}')]))
schema = {"type": "object", "properties": {"text": {"type": "string"}}}
response = await client.get_response(
messages=[Message(role="user", contents=[Content.from_text("Hi")])],
options={"response_format": schema},
)
assert response.value == {"text": "hello"}
async def test_response_schema_added_to_config() -> None:
"""Sets both response_mime_type and the raw schema on the config when response_schema is given."""
client, mock = _make_gemini_client()
@@ -931,6 +945,284 @@ async def test_response_schema_added_to_config() -> None:
assert config.response_schema == schema
async def test_response_format_raw_json_schema_added_to_config() -> None:
"""For declarative outputSchema, response_format may already be a raw JSON schema mapping."""
client, mock = _make_gemini_client()
mock.aio.models.generate_content = AsyncMock(return_value=_make_response([_make_part(text='{"answer": "hello"}')]))
schema = {
"type": "object",
"properties": {"answer": {"type": "string", "description": "The answer."}},
"required": ["answer"],
}
await client.get_response(
messages=[Message(role="user", contents=[Content.from_text("Hi")])],
options={"response_format": schema},
)
config: types.GenerateContentConfig = mock.aio.models.generate_content.call_args.kwargs["config"]
assert config.response_mime_type == "application/json"
assert config.response_schema == schema
async def test_agent_default_options_response_format_raw_schema_added_to_config() -> None:
"""Agent default_options is the path used by declarative outputSchema."""
client, mock = _make_gemini_client()
mock.aio.models.generate_content = AsyncMock(return_value=_make_response([_make_part(text='{"answer": "hello"}')]))
schema = {"type": "object", "properties": {"answer": {"type": "string"}}, "required": ["answer"]}
agent = Agent(client=client, default_options={"response_format": schema})
await agent.run("Hi")
config: types.GenerateContentConfig = mock.aio.models.generate_content.call_args.kwargs["config"]
assert config.response_mime_type == "application/json"
assert config.response_schema == schema
async def test_response_format_complex_raw_json_schema_preserved() -> None:
"""Nested declarative schemas should be forwarded without losing shape or constraints."""
client, mock = _make_gemini_client()
mock.aio.models.generate_content = AsyncMock(return_value=_make_response([_make_part(text='{"answer": "ok"}')]))
schema = {
"type": "object",
"properties": {
"answer": {"type": "string"},
"citations": {
"type": "array",
"items": {
"type": "object",
"properties": {
"source": {"type": "string"},
"confidence": {"type": "number"},
},
"required": ["source"],
"additionalProperties": False,
},
},
},
"required": ["answer"],
"additionalProperties": False,
}
await client.get_response(
messages=[
Message(
role="user",
contents=[
Content.from_text(
"Summarize a long document while preserving citation metadata.\n" + ("context\n" * 128)
)
],
)
],
options={"response_format": schema},
)
config: types.GenerateContentConfig = mock.aio.models.generate_content.call_args.kwargs["config"]
assert config.response_schema == schema
async def test_response_format_json_schema_envelope_added_to_config() -> None:
"""OpenAI-style json_schema envelopes should still provide Gemini with the inner schema."""
client, mock = _make_gemini_client()
mock.aio.models.generate_content = AsyncMock(return_value=_make_response([_make_part(text='{"answer": "hello"}')]))
schema = {"type": "object", "properties": {"answer": {"type": "string"}}}
await client.get_response(
messages=[Message(role="user", contents=[Content.from_text("Hi")])],
options={"response_format": {"type": "json_schema", "json_schema": {"name": "Answer", "schema": schema}}},
)
config: types.GenerateContentConfig = mock.aio.models.generate_content.call_args.kwargs["config"]
assert config.response_mime_type == "application/json"
assert config.response_schema == schema
async def test_response_format_format_envelope_added_to_config() -> None:
"""Responses-style format envelopes should also provide Gemini with the nested schema."""
client, mock = _make_gemini_client()
mock.aio.models.generate_content = AsyncMock(return_value=_make_response([_make_part(text='{"answer": "hello"}')]))
schema = {"type": "object", "properties": {"answer": {"type": "string"}}}
await client.get_response(
messages=[Message(role="user", contents=[Content.from_text("Hi")])],
options={"response_format": {"format": {"type": "json_schema", "name": "Answer", "schema": schema}}},
)
config: types.GenerateContentConfig = mock.aio.models.generate_content.call_args.kwargs["config"]
assert config.response_mime_type == "application/json"
assert config.response_schema == schema
async def test_response_format_direct_schema_key_added_to_config() -> None:
"""Provider-normalized mappings with a direct schema key should be accepted."""
client, mock = _make_gemini_client()
mock.aio.models.generate_content = AsyncMock(return_value=_make_response([_make_part(text='{"answer": "hello"}')]))
schema = {"type": "object", "properties": {"answer": {"type": "string"}}}
await client.get_response(
messages=[Message(role="user", contents=[Content.from_text("Hi")])],
options={"response_format": {"schema": schema}},
)
config: types.GenerateContentConfig = mock.aio.models.generate_content.call_args.kwargs["config"]
assert config.response_mime_type == "application/json"
assert config.response_schema == schema
async def test_response_format_json_schema_envelope_preserves_empty_schema() -> None:
"""An explicitly empty JSON schema is still a schema and should not be dropped as falsy."""
client, mock = _make_gemini_client()
mock.aio.models.generate_content = AsyncMock(return_value=_make_response([_make_part(text="{}")]))
schema: dict[str, Any] = {}
await client.get_response(
messages=[Message(role="user", contents=[Content.from_text("Hi")])],
options={"response_format": {"type": "json_schema", "json_schema": {"name": "AnyJson", "schema": schema}}},
)
config: types.GenerateContentConfig = mock.aio.models.generate_content.call_args.kwargs["config"]
assert config.response_schema == schema
async def test_response_format_anyof_raw_schema_added_to_config() -> None:
"""Raw schemas without a type should still be recognized when they use JSON Schema keywords."""
client, mock = _make_gemini_client()
mock.aio.models.generate_content = AsyncMock(return_value=_make_response([_make_part(text='"ok"')]))
schema = {"anyOf": [{"type": "string"}, {"type": "number"}]}
await client.get_response(
messages=[Message(role="user", contents=[Content.from_text("Hi")])],
options={"response_format": schema},
)
config: types.GenerateContentConfig = mock.aio.models.generate_content.call_args.kwargs["config"]
assert config.response_schema == schema
async def test_response_format_union_type_raw_schema_added_to_config() -> None:
"""JSON Schema union type arrays should be treated as raw schemas."""
client, mock = _make_gemini_client()
mock.aio.models.generate_content = AsyncMock(return_value=_make_response([_make_part(text='{"answer": "hello"}')]))
schema = {"type": ["object", "null"], "properties": {"answer": {"type": "string"}}}
await client.get_response(
messages=[Message(role="user", contents=[Content.from_text("Hi")])],
options={"response_format": schema},
)
config: types.GenerateContentConfig = mock.aio.models.generate_content.call_args.kwargs["config"]
assert config.response_schema == schema
async def test_response_format_json_object_does_not_set_schema() -> None:
"""A JSON-object response_format requests JSON output but is not itself a Gemini response schema."""
client, mock = _make_gemini_client()
mock.aio.models.generate_content = AsyncMock(return_value=_make_response([_make_part(text="{}")]))
await client.get_response(
messages=[Message(role="user", contents=[Content.from_text("Hi")])],
options={"response_format": {"type": "json_object"}},
)
config: types.GenerateContentConfig = mock.aio.models.generate_content.call_args.kwargs["config"]
assert config.response_mime_type == "application/json"
assert config.response_schema is None
async def test_response_format_json_schema_without_inner_schema_does_not_set_schema() -> None:
"""A json_schema envelope without a schema should not be mistaken for a raw JSON schema."""
client, mock = _make_gemini_client()
mock.aio.models.generate_content = AsyncMock(return_value=_make_response([_make_part(text="{}")]))
await client.get_response(
messages=[Message(role="user", contents=[Content.from_text("Hi")])],
options={"response_format": {"type": "json_schema", "json_schema": {"name": "MissingSchema"}}},
)
config: types.GenerateContentConfig = mock.aio.models.generate_content.call_args.kwargs["config"]
assert config.response_mime_type == "application/json"
assert config.response_schema is None
async def test_response_schema_takes_precedence_over_response_format_schema() -> None:
"""An explicit Gemini response_schema should win when both schema options are present."""
client, mock = _make_gemini_client()
mock.aio.models.generate_content = AsyncMock(return_value=_make_response([_make_part(text="{}")]))
response_format_schema = {"type": "object", "properties": {"name": {"type": "string"}}}
response_schema = {"type": "object", "properties": {"id": {"type": "integer"}}}
await client.get_response(
messages=[Message(role="user", contents=[Content.from_text("Hi")])],
options={"response_format": response_format_schema, "response_schema": response_schema},
)
config: types.GenerateContentConfig = mock.aio.models.generate_content.call_args.kwargs["config"]
assert config.response_schema == response_schema
async def test_response_format_raw_schema_kept_with_tools() -> None:
"""Structured output must still reach Gemini when function tools are present."""
def calculator(expression: str) -> str:
"""Evaluate a simple expression."""
return expression
tool = FunctionTool(name="calculator", func=calculator)
client, mock = _make_gemini_client()
mock.aio.models.generate_content = AsyncMock(return_value=_make_response([_make_part(text='{"answer": "4"}')]))
schema = {"type": "object", "properties": {"answer": {"type": "string"}}, "required": ["answer"]}
await client.get_response(
messages=[Message(role="user", contents=[Content.from_text("What is 2 + 2?")])],
options={"tools": [tool], "response_format": schema},
)
config: types.GenerateContentConfig = mock.aio.models.generate_content.call_args.kwargs["config"]
assert config.response_schema == schema
assert config.tools is not None
assert config.tools[0].function_declarations[0].name == "calculator"
async def test_streaming_response_format_raw_schema_added_to_config() -> None:
"""Streaming requests use the same config path and should also forward raw schema mappings."""
client, mock = _make_gemini_client()
chunks = [_make_response([_make_part(text='{"answer": "hello"}')], finish_reason="STOP")]
mock.aio.models.generate_content_stream = AsyncMock(return_value=_async_iter(chunks))
schema = {"type": "object", "properties": {"answer": {"type": "string"}}}
stream = client.get_response(
messages=[Message(role="user", contents=[Content.from_text("Hi")])],
options={"response_format": schema},
stream=True,
)
async for _ in stream:
pass
config: types.GenerateContentConfig = mock.aio.models.generate_content_stream.call_args.kwargs["config"]
assert config.response_mime_type == "application/json"
assert config.response_schema == schema
async def test_streaming_response_format_mapping_populates_final_value() -> None:
"""Streaming responses should preserve mapping response_format for final value parsing."""
client, mock = _make_gemini_client()
chunks = [_make_response([_make_part(text='{"answer": "hello"}')], finish_reason="STOP")]
mock.aio.models.generate_content_stream = AsyncMock(return_value=_async_iter(chunks))
schema = {"type": "object", "properties": {"answer": {"type": "string"}}}
stream = client.get_response(
messages=[Message(role="user", contents=[Content.from_text("Hi")])],
options={"response_format": schema},
stream=True,
)
async for _ in stream:
pass
final = await stream.get_final_response()
assert final.value == {"answer": "hello"}
async def test_streaming_response_format_passed_to_build_response_stream() -> None:
"""Verifies that response_format is forwarded to _build_response_stream when streaming
so that structured output parsing works correctly on the final assembled response.
+8
View File
@@ -13,9 +13,12 @@ The Model Context Protocol (MCP) is an open standard for connecting AI agents to
| **Agent as MCP Server** | [`agent_as_mcp_server.py`](agent_as_mcp_server.py) | Shows how to expose an Agent Framework agent as an MCP server that other AI applications can connect to |
| **API Key Authentication** | [`mcp_api_key_auth.py`](mcp_api_key_auth.py) | Demonstrates API key authentication with MCP servers using `header_provider`, runtime invocation kwargs, and a command-line API key argument |
| **GitHub Integration with PAT** | [`mcp_github_pat.py`](mcp_github_pat.py) | Demonstrates connecting to GitHub's MCP server using Personal Access Token (PAT) authentication |
| **Long-Running Task** | [`mcp_long_running_task.py`](mcp_long_running_task.py) | Demonstrates transparent SEP-2663 long-running task handling for MCP tools that advertise `taskSupport=required`. Self-spawns a stdio MCP child server |
## Prerequisites
Most samples in this folder use OpenAI:
- `OPENAI_API_KEY` environment variable
- `OPENAI_CHAT_MODEL` environment variable
@@ -23,3 +26,8 @@ Run `mcp_api_key_auth.py` with the MCP API key as the first command-line argumen
For `mcp_github_pat.py`:
- `GITHUB_PAT` - Your GitHub Personal Access Token (create at https://github.com/settings/tokens)
For `mcp_long_running_task.py` (uses Azure OpenAI via Entra-ID):
- Run `az login` once
- `AZURE_OPENAI_ENDPOINT` - your Azure OpenAI resource endpoint, e.g. `https://<resource>.openai.azure.com/`
- `AZURE_OPENAI_CHAT_MODEL` (or `AZURE_OPENAI_MODEL`) - the deployment name (e.g. `gpt-4o-mini`)
@@ -46,6 +46,8 @@ async def github_mcp_example() -> None:
# The MCP tool manages the connection to the MCP server and makes its tools available
# Set approval_mode="never_require" to allow the MCP tool to execute without approval
client = OpenAIChatClient()
# Note that the tool created here will be executed remotely by OpenAI, not locally by
# your application.
github_mcp_tool = client.get_mcp_tool(
name="GitHub",
url="https://api.githubcopilot.com/mcp/",
@@ -0,0 +1,181 @@
# Copyright (c) Microsoft. All rights reserved.
"""
MCP Long-Running Task (SEP-2663) Example
Demonstrates that ``MCPStdioTool`` transparently drives the MCP long-running
task lifecycle for tools that advertise ``execution.taskSupport == "required"``.
The agent observes a single function-call result; the framework handles the
``tools/call`` → ``tasks/get`` (polled) → ``tasks/result`` sequence in the
background.
Run it as a single file. The script doubles as both the client and the stdio
MCP child server (the child branch is selected via ``--server``):
python mcp_long_running_task.py
Requirements:
- Azure CLI sign-in (``az login``) — used for Entra-ID auth against Azure OpenAI.
- ``AZURE_OPENAI_ENDPOINT`` — your Azure OpenAI resource endpoint, e.g.
``https://<resource>.openai.azure.com/``.
- ``AZURE_OPENAI_CHAT_MODEL`` (or ``AZURE_OPENAI_MODEL``) — the deployment name,
e.g. ``gpt-4o-mini``.
This sample uses the lower-level ``mcp.server.lowlevel.Server`` so it can:
1. Advertise a tool with ``execution=ToolExecution(taskSupport="required")``.
2. Enable the SDK's experimental task support for the ``tasks/*`` lifecycle.
"""
import asyncio
import sys
from datetime import timedelta
from typing import Any
from agent_framework import Agent, MCPStdioTool, MCPTaskOptions
from agent_framework.openai import OpenAIChatClient
from azure.identity import AzureCliCredential
from dotenv import load_dotenv
load_dotenv()
# ---------------------------------------------------------------------------
# MCP stdio server (child-process branch)
# ---------------------------------------------------------------------------
async def _run_server() -> None:
"""Run a minimal stdio MCP server exposing one long-running tool."""
import mcp.types as types
from mcp.server.lowlevel import Server
from mcp.server.stdio import stdio_server
server: Server[Any, Any] = Server("mcp-long-running-task-demo")
# Auto-registers handlers for tasks/get, tasks/result, tasks/cancel, tasks/list
# backed by an in-memory store.
server.experimental.enable_tasks()
@server.list_tools()
async def _list_tools() -> list[types.Tool]: # pyright: ignore[reportUnusedFunction]
return [
types.Tool(
name="slow_summary",
description=(
"Produces a short summary of the supplied text after simulating several seconds of expensive work."
),
inputSchema={
"type": "object",
"properties": {
"text": {
"type": "string",
"description": "Text to summarize.",
}
},
"required": ["text"],
},
# Advertise that this tool MUST be invoked via the task lifecycle.
execution=types.ToolExecution(taskSupport="required"),
)
]
@server.call_tool()
async def _call_tool(name: str, arguments: dict[str, Any]) -> Any: # pyright: ignore[reportUnusedFunction]
if name != "slow_summary":
raise ValueError(f"Unknown tool: {name}")
ctx = server.request_context
async def _work(task: Any) -> types.CallToolResult:
await task.update_status("Thinking...")
await asyncio.sleep(15.0)
text: str = (arguments.get("text") or "").strip()
words = text.split()
preview = " ".join(words[:6]) + ("..." if len(words) > 6 else "")
summary = (
f"Summarized {len(words)} word(s). First few words: '{preview}'."
if words
else "No input text was provided."
)
return types.CallToolResult(
content=[types.TextContent(type="text", text=summary)],
isError=False,
)
if not ctx.experimental.is_task:
# Client invoked the tool without task augmentation. Return a hard
# error so a misconfigured client surfaces the problem clearly.
return types.CallToolResult(
content=[
types.TextContent(
type="text",
text="'slow_summary' must be invoked as a task.",
)
],
isError=True,
)
return await ctx.experimental.run_task(_work)
async with stdio_server() as (read_stream, write_stream):
await server.run(read_stream, write_stream, server.create_initialization_options())
# ---------------------------------------------------------------------------
# Agent client (default branch)
# ---------------------------------------------------------------------------
async def _run_client() -> None:
mcp_tool = MCPStdioTool(
name="LongRunningDemo",
description="Demo MCP server exposing a tool that advertises taskSupport=required.",
command=sys.executable,
args=[__file__, "--server"],
# Optional: cap individual tasks at two minutes. The server may apply its
# own default if this is omitted.
task_options=MCPTaskOptions(default_ttl=timedelta(minutes=2)),
)
async with Agent(
client=OpenAIChatClient(credential=AzureCliCredential()),
name="LROAgent",
instructions=(
"You are a helpful assistant. Use the slow_summary tool when the user "
"asks for a summary. Wait for the result and present it directly."
),
tools=mcp_tool,
) as agent:
prompt = (
"Please summarize the following text using your slow_summary tool: "
"'The Model Context Protocol lets language models talk to external "
"tools and resources through a small JSON-RPC surface.'"
)
print("=== run() ===")
print(f"User: {prompt}")
response = await agent.run(prompt)
print(f"Agent: {response.text}\n")
print("=== run(stream=True) ===")
print(f"User: {prompt}")
print("Agent: ", end="", flush=True)
async for update in agent.run(prompt, stream=True):
if update.text:
print(update.text, end="", flush=True)
print()
# ---------------------------------------------------------------------------
# Entry point
# ---------------------------------------------------------------------------
def main() -> None:
if len(sys.argv) > 1 and sys.argv[1] == "--server":
asyncio.run(_run_server())
return
asyncio.run(_run_client())
if __name__ == "__main__":
main()
@@ -134,15 +134,11 @@ def handle_response_and_requests(response: AgentResponse) -> dict[str, HandoffAg
if message.text:
print(f"- {message.author_name or message.role}: {message.text}")
for content in message.contents:
if content.type == "function_call":
if isinstance(content.arguments, dict):
request = WorkflowAgent.RequestInfoFunctionArgs.from_dict(content.arguments)
elif isinstance(content.arguments, str):
request = WorkflowAgent.RequestInfoFunctionArgs.from_json(content.arguments)
else:
raise ValueError("Invalid arguments type. Expecting a request info structure for this sample.")
if isinstance(request.data, HandoffAgentUserRequest):
pending_requests[request.request_id] = request.data
if content.type == "function_call" and content.name == WorkflowAgent.REQUEST_INFO_FUNCTION_NAME:
request_function_args = WorkflowAgent.RequestInfoFunctionArgs.from_dict(content.arguments) # type: ignore
request_id = request_function_args.request_id
request_event = request_function_args.request_event
pending_requests[request_id] = request_event.data
return pending_requests
@@ -3,10 +3,8 @@
import asyncio
import os
import sys
from collections.abc import Mapping
from dataclasses import dataclass
from pathlib import Path
from typing import Any
from agent_framework.foundry import FoundryChatClient
from azure.identity import AzureCliCredential
@@ -141,28 +139,14 @@ async def main() -> None:
# Handle the human review if required.
if human_review_function_call:
# Parse the human review request arguments.
human_request_args = human_review_function_call.arguments
if isinstance(human_request_args, str):
request: WorkflowAgent.RequestInfoFunctionArgs = WorkflowAgent.RequestInfoFunctionArgs.from_json(
human_request_args
)
elif isinstance(human_request_args, Mapping):
request = WorkflowAgent.RequestInfoFunctionArgs.from_dict(dict(human_request_args))
else:
raise TypeError("Unexpected argument type for human review function call.")
request_payload: Any = request.data
human_request_args = WorkflowAgent.RequestInfoFunctionArgs.from_dict(human_review_function_call.arguments) # type: ignore
request_payload = human_request_args.request_event.data
if not isinstance(request_payload, HumanReviewRequest):
raise ValueError("Human review request payload must be a HumanReviewRequest.")
agent_request = request_payload.agent_request
if agent_request is None:
raise ValueError("Human review request must include agent_request.")
request_id = agent_request.request_id
if not request_payload.agent_request:
raise ValueError("Human review request must contain an agent_request.")
# Mock a human response approval for demonstration purposes.
human_response = ReviewResponse(request_id=request_id, feedback="", approved=True)
human_response = ReviewResponse(request_id=request_payload.agent_request.request_id, feedback="", approved=True)
# Create the function call result object to send back to the agent.
human_review_function_result = Content(
"function_result",
@@ -28,6 +28,7 @@ import zipfile
from pathlib import Path
from azure.ai.projects.aio import AIProjectClient
from azure.ai.projects.models import CreateSkillVersionFromFilesBody
from azure.core.exceptions import ResourceNotFoundError
from azure.identity.aio import DefaultAzureCredential
from dotenv import load_dotenv
@@ -68,8 +69,13 @@ async def main() -> None:
name = skill_md.parent.name
print(f"Provisioning skill '{name}' from {skill_md.relative_to(SKILLS_DIR.parent)}...")
await _delete_skill_if_exists(project, name)
imported = await project.beta.skills.create_from_package(_zip_skill_md(skill_md))
print(f" Imported skill '{imported.name}' (id={imported.skill_id}, has_blob={imported.has_blob}).")
imported = await project.beta.skills.create_from_files(
name,
content=CreateSkillVersionFromFilesBody(
files=[(f"{name}.zip", _zip_skill_md(skill_md), "application/zip")]
),
)
print(f" Imported skill '{imported.name}' (id={imported.skill_id}, version={imported.version}).")
print("Verifying skills via project.beta.skills.list()...")
listed = {skill.name: skill async for skill in project.beta.skills.list()}
@@ -79,8 +85,8 @@ async def main() -> None:
if skill is None:
raise RuntimeError(f"Skill '{name}' was imported but is not present in the project listing.")
print(
f" OK '{skill.name}': id={skill.skill_id}, "
f"description={skill.description!r}, has_blob={skill.has_blob}"
f" OK '{skill.name}': id={skill.id}, "
f"description={skill.description!r}, default_version={skill.default_version}"
)
print("Done.")