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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
41 changes: 20 additions & 21 deletions tools/azsdk-cli/Azure.Sdk.Tools.Cli.Contract/MCPTool.cs
Original file line number Diff line number Diff line change
Expand Up @@ -3,32 +3,31 @@
using System.CommandLine;
using System.CommandLine.Invocation;

namespace Azure.Sdk.Tools.Cli.Contract
namespace Azure.Sdk.Tools.Cli.Contract;

/// <summary>
/// This is the base class defining how an MCP enabled tool will interface with the server.
///
/// This covers:
/// - route registration/disambiguation
/// - compilation trim avoidance for reflection-included MCP tools
/// </summary>
public abstract class MCPTool
{
/// <summary>
/// This is the base class defining how an MCP enabled tool will interface with the server.
///
/// This covers:
/// - route registration/disambiguation
/// - compilation trim avoidance for reflection-included MCP tools
/// </summary>
public abstract class MCPTool
{
public MCPTool() { }
public MCPTool() { }

public Command? Command;
public Command? Command;

public int ExitCode { get; set; } = 0;
public int ExitCode { get; set; } = 0;

public void SetFailure(int exitCode = 1)
{
ExitCode = exitCode;
}
public void SetFailure(int exitCode = 1)
{
ExitCode = exitCode;
}

public CommandGroup[] CommandHierarchy { get; set; } = [];
public CommandGroup[] CommandHierarchy { get; set; } = [];

public abstract Command GetCommand();
public abstract Command GetCommand();

public abstract Task HandleCommand(InvocationContext ctx, CancellationToken ct);
}
public abstract Task HandleCommand(InvocationContext ctx, CancellationToken ct);
}
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,8 @@
using Azure.AI.OpenAI;
using Azure.Sdk.Tools.Cli.Microagents;
using Azure.Sdk.Tools.Cli.Microagents.Tools;
using Azure.Sdk.Tools.Cli.Helpers;
using NUnit.Framework;
using Moq;
using OpenAI.Chat;

Expand All @@ -23,7 +25,8 @@ public void Setup()
chatClientMock = new Mock<ChatClient>();
openAIClientMock.Setup(client => client.GetChatClient(It.IsAny<string>()))
.Returns(chatClientMock.Object);
microagentHostService = new MicroagentHostService(openAIClientMock.Object, loggerMock.Object);
var tokenUsageHelper = new TokenUsageHelper(Mock.Of<Azure.Sdk.Tools.Cli.Helpers.IOutputHelper>());
microagentHostService = new MicroagentHostService(openAIClientMock.Object, loggerMock.Object, tokenUsageHelper);
}

[Test]
Expand Down
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
// Copyright (c) Microsoft Corporation.
// Licensed under the MIT License.
using Moq;
using NUnit.Framework;
using NUnit.Framework.Internal;
using Azure.Core;
using Azure.Sdk.Tools.Cli.Helpers;
Expand All @@ -22,6 +23,7 @@ internal class ExampleToolTests
private MockGitHubService? mockGitHubService;
private Mock<IProcessHelper>? mockProcessHelper;
private Mock<IPowershellHelper>? mockPowershellHelper;
private Mock<Azure.Sdk.Tools.Cli.Microagents.IMicroagentHostService>? mockMicroagentHostService;

[SetUp]
public void Setup()
Expand All @@ -33,6 +35,7 @@ public void Setup()
mockGitHubService = new MockGitHubService();
mockProcessHelper = new Mock<IProcessHelper>();
mockPowershellHelper = new Mock<IPowershellHelper>();
mockMicroagentHostService = new Mock<Azure.Sdk.Tools.Cli.Microagents.IMicroagentHostService>();

// Set up Azure service mock to return a mock credential
var mockCredential = new Mock<TokenCredential>();
Expand All @@ -58,6 +61,8 @@ public void Setup()
mockGitHubService,
mockProcessHelper.Object,
mockPowershellHelper.Object,
tokenUsageHelper: new TokenUsageHelper(mockOutput.Object),
mockMicroagentHostService.Object,
#pragma warning disable CS8625 // Cannot convert null literal to non-nullable reference type.
null
#pragma warning restore CS8625 // Cannot convert null literal to non-nullable reference type.
Expand Down Expand Up @@ -153,7 +158,7 @@ public void GetCommand_ReturnsCommandWithCorrectSubCommands()

Assert.That(command.Name, Is.EqualTo("demo"));
Assert.That(command.Description, Does.Contain("Comprehensive demonstration"));
Assert.That(command.Subcommands.Count, Is.EqualTo(7));
Assert.That(command.Subcommands.Count, Is.EqualTo(8));

var subCommandNames = command.Subcommands.Select(sc => sc.Name).ToList();
Assert.That(subCommandNames, Does.Contain("azure"));
Expand All @@ -163,6 +168,7 @@ public void GetCommand_ReturnsCommandWithCorrectSubCommands()
Assert.That(subCommandNames, Does.Contain("error"));
Assert.That(subCommandNames, Does.Contain("process"));
Assert.That(subCommandNames, Does.Contain("powershell"));
Assert.That(subCommandNames, Does.Contain("microagent"));
}

[Test]
Expand Down Expand Up @@ -236,7 +242,6 @@ public async Task DemonstratePowershellExecution_Success()
{
mockPowershellHelper!.Setup(p => p.Run(It.IsAny<PowershellOptions>(), It.IsAny<CancellationToken>()))
.ReturnsAsync(new ProcessResult { ExitCode = 0, });

var result = await tool.DemonstratePowershellExecution("foobar");

Assert.That(result.ResponseError, Is.Null);
Expand All @@ -245,4 +250,16 @@ public async Task DemonstratePowershellExecution_Success()
Assert.That(result.Result, Is.Empty);
Assert.That(result.Details?["exit_code"], Is.EqualTo("0"));
}

[Test]
public async Task DemonstrateMicroagentFibonacci_Success()
{
mockMicroagentHostService!.Setup(m => m.RunAgentToCompletion(It.IsAny<Azure.Sdk.Tools.Cli.Microagents.Microagent<int>>(), It.IsAny<CancellationToken>()))
.ReturnsAsync(13);

var response = await tool.DemonstrateMicroagentFibonacci(7);

Assert.That(response.ResponseError, Is.Null);
Assert.That(response.Result as string, Does.Contain("Fibonacci(7) = 13"));
}
}
84 changes: 14 additions & 70 deletions tools/azsdk-cli/Azure.Sdk.Tools.Cli/Helpers/TokenUsageHelper.cs
Original file line number Diff line number Diff line change
@@ -1,80 +1,24 @@
namespace Azure.Sdk.Tools.Cli.Helpers;

public class TokenUsageHelper
public class TokenUsageHelper(IOutputHelper outputHelper)
{
protected double PromptTokens { get; set; }
protected double CompletionTokens { get; set; }
protected double InputCost { get; set; }
protected double OutputCost { get; set; }
protected double TotalCost { get; set; }
public List<string> Models { get; set; } = [];
protected double PromptTokens { get; set; } = 0;
protected double CompletionTokens { get; set; } = 0;
protected IEnumerable<string> ModelsUsed { get; set; } = [];

public TokenUsageHelper(string model, long inputTokens, long outputTokens)
public void Add(string model, long inputTokens, long outputTokens)
{
PromptTokens = inputTokens;
CompletionTokens = outputTokens;
Models = [model];
SetCost(model);
ModelsUsed = ModelsUsed.Union([model]);
PromptTokens += inputTokens;
CompletionTokens += outputTokens;
}

protected TokenUsageHelper() { }

private void SetCost(string model)
public void LogUsage()
{
var oneMillion = 1000000;
double inputPrice, outputPrice;

// Prices assume the slightly more expensive regional model pricing
if (model == "gpt-4o")
{
(inputPrice, outputPrice) = (2.75, 11);
}
else if (model == "gpt-4o-mini")
{
(inputPrice, outputPrice) = (0.165, 0.66);
}
if (model == "gpt-4.1")
{
(inputPrice, outputPrice) = (2, 8);
}
else if (model == "gpt-4.1-mini")
{
(inputPrice, outputPrice) = (0.4, 1.60);
}
else if (model == "o3-mini")
{
(inputPrice, outputPrice) = (1.21, 4.84);
}
else
{
return;
}


InputCost = PromptTokens / oneMillion * inputPrice;
OutputCost = CompletionTokens / oneMillion * outputPrice;
}
var models = string.Join(", ", ModelsUsed);

public void LogCost()
{
var _inputCost = InputCost == 0 ? "?" : InputCost.ToString("F3");
var _outputCost = OutputCost == 0 ? "?" : OutputCost.ToString("F3");
var _totalCost = (InputCost + OutputCost) == 0 ? "?" : (InputCost + OutputCost).ToString("F3");
var models = string.Join(", ", Models);
Console.WriteLine("--------------------------------------------------------------------------------");
Console.WriteLine($"[{models}] Usage (cost / tokens):");
Console.WriteLine($" Input: ${_inputCost} / {PromptTokens}");
Console.WriteLine($" Output: ${_outputCost} / {CompletionTokens}");
Console.WriteLine($" Total: ${_totalCost} / {PromptTokens + CompletionTokens}");
Console.WriteLine("--------------------------------------------------------------------------------");
outputHelper.OutputConsole("--------------------------------------------------------------------------------");
outputHelper.OutputConsole($"[token usage][{models}] input: {PromptTokens}, output: {CompletionTokens}, total: {PromptTokens + CompletionTokens}");
outputHelper.OutputConsole("--------------------------------------------------------------------------------");
}

public static TokenUsageHelper operator +(TokenUsageHelper a, TokenUsageHelper? b) => new()
{
Models = a.Models.Union(b?.Models ?? []).ToList(),
PromptTokens = a.PromptTokens + (b?.PromptTokens ?? 0),
CompletionTokens = a.CompletionTokens + (b?.CompletionTokens ?? 0),
InputCost = a.InputCost + (b?.InputCost ?? 0),
OutputCost = a.OutputCost + (b?.OutputCost ?? 0),
};
}
}
Original file line number Diff line number Diff line change
@@ -1,10 +1,11 @@
using System.ComponentModel;
using Azure.AI.OpenAI;
using Azure.Sdk.Tools.Cli.Helpers;
using OpenAI.Chat;

namespace Azure.Sdk.Tools.Cli.Microagents;

public class MicroagentHostService(AzureOpenAIClient openAI, ILogger<MicroagentHostService> logger) : IMicroagentHostService
public class MicroagentHostService(AzureOpenAIClient openAI, ILogger<MicroagentHostService> logger, TokenUsageHelper tokenUsageHelper) : IMicroagentHostService
{
private const string ExitToolName = "Exit";

Expand Down Expand Up @@ -60,6 +61,10 @@ public async Task<TResult> RunAgentToCompletion<TResult>(Microagent<TResult> age
// Request the chat completion
logger.LogDebug("Sending conversation history with {MessageCount} messages to model '{Model}'", conversationHistory.Count, agentDefinition.Model);
var response = await chatClient.CompleteChatAsync(conversationHistory, chatCompletionOptions, ct);
if (null != response.Value.Usage)
{
tokenUsageHelper.Add(agentDefinition.Model, response.Value.Usage.InputTokenCount, response.Value.Usage.OutputTokenCount);
}

var toolCall = response.Value.ToolCalls.Single();
logger.LogInformation("Model called tool '{ToolName}'", toolCall.FunctionName);
Expand Down Expand Up @@ -104,7 +109,7 @@ public async Task<TResult> RunAgentToCompletion<TResult>(Microagent<TResult> age
conversationHistory.Add(ChatMessage.CreateToolMessage(toolCall.Id, toolResult));
}

throw new Exception("Agent did not return a result within the maximum number of iterations");
throw new Exception($"Agent did not return a result within the maximum number of {agentDefinition.MaxToolCalls} iterations");
}

/// <summary>
Expand Down
22 changes: 10 additions & 12 deletions tools/azsdk-cli/Azure.Sdk.Tools.Cli/Program.cs
Original file line number Diff line number Diff line change
Expand Up @@ -46,7 +46,8 @@ public static WebApplicationBuilder CreateAppBuilder(string[] args)
WebApplicationBuilder builder = WebApplication.CreateBuilder(args);

builder.Services.AddOpenTelemetry()
.WithTracing(b => {
.WithTracing(b =>
{
b.AddSource(Constants.TOOLS_ACTIVITY_SOURCE)
.AddAspNetCoreInstrumentation()
.AddHttpClientInstrumentation()
Expand All @@ -55,24 +56,21 @@ public static WebApplicationBuilder CreateAppBuilder(string[] args)
})
.UseOtlpExporter();

// Log everything to stderr in mcp mode so the client doesn't try to interpret stdout messages that aren't json rpc
var logErrorThreshold = isCLI ? LogLevel.Error : LogLevel.Debug;

builder.Logging.AddConsole(consoleLogOptions =>
{
// Log everything to stderr in mcp mode so the client doesn't try to interpret stdout messages that aren't json rpc
var logErrorThreshold = isCLI ? LogLevel.Error : LogLevel.Debug;
consoleLogOptions.LogToStandardErrorThreshold = logErrorThreshold;
});

// Skip verbose azure client logging
// Skip azure client logging noise
builder.Logging.AddFilter((category, level) =>
{
var isAzureClient = category!.StartsWith("Azure.", StringComparison.Ordinal);
var isToolsClient = category!.StartsWith("Azure.Sdk.Tools.", StringComparison.Ordinal);
if (isAzureClient && !isToolsClient)
{
return level >= LogLevel.Warning;
}
return level >= logErrorThreshold;
if (debug || null == category) { return level >= logLevel; }
var isAzureClient = category.StartsWith("Azure.", StringComparison.Ordinal);
var isToolsClient = category.StartsWith("Azure.Sdk.Tools.", StringComparison.Ordinal);
if (isAzureClient && !isToolsClient) { return level >= LogLevel.Error; }
return level >= logLevel;
});

// add the console logger
Expand Down
Loading
Loading