Skip to content

Commit a3dd028

Browse files
fix: address PR #1073 review feedback (security hardening iteration 3)
- ResponseIdValidationService: redesign from per-user HashSet to per-responseId cache entries, eliminating arbitrary eviction (ids.Remove(ids.First())) and reducing lock contention; each ID now has its own 2-hour TTL - McpApiTokenService.CreateTokenAsync: enforce MaxTokensPerUser limit inside a serializable transaction, making the cap atomic under concurrency; move limit check out of controller into service - McpTokenController: remove pre-check (now redundant); catch InvalidOperationException from service and return 400 - chat-module.js: clear sessionStorage lastResponseId when localStorage history is rejected due to TTL expiry, preventing stale responseId - AIChatService: add LogMcpToolCallInvokedStream [LoggerMessage] for the streaming code path (previously shared LogMcpToolCallInvoked with iteration semantics, now split: iteration vs depth are distinct) - AIChatService: reword endUserId comment to accurately reflect that the SDK does not support it yet (was misleadingly claiming it was forwarded) - Tests: add GetActiveTokenCountAsync tests (active, revoked, expired, zero, at-max-limit) to McpApiTokenServiceTests - Tests: add ResponseIdValidationServiceTests (new conversation, cache miss, owner validates, cross-user rejected, null inputs, multi-user isolation)
1 parent b1119cc commit a3dd028

7 files changed

Lines changed: 275 additions & 61 deletions

File tree

EssentialCSharp.Chat.Shared/Services/AIChatService.cs

Lines changed: 9 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -264,7 +264,7 @@ private static string SanitizeForXmlContext(string? input) =>
264264
yield break;
265265
}
266266

267-
LogMcpToolCallInvoked(_Logger, functionCallItem.FunctionName, toolCallDepth);
267+
LogMcpToolCallInvokedStream(_Logger, functionCallItem.FunctionName, toolCallDepth);
268268
// A dictionary of arguments to pass to the tool. Each key represents a parameter name, and its associated value represents the argument value.
269269
Dictionary<string, object?> arguments = [];
270270
// example JsonResponse:
@@ -343,12 +343,12 @@ private async Task<ResponseCreationOptions> CreateResponseOptionsAsync(
343343
options.PreviousResponseId = previousResponseId;
344344
}
345345

346-
// Forward the authenticated end-user's identifier to Azure OpenAI.
347-
// This enables Microsoft Defender for Cloud's prompt-shield and abuse detection.
346+
// endUserId is reserved for forwarding to Azure OpenAI for end-user attribution
347+
// (Microsoft Defender prompt-shield correlation). OpenAI .NET SDK v2.7.0 does not
348+
// expose ResponseCreationOptions.User; this parameter is intentionally discarded
349+
// until SDK support is available.
348350
// See: https://learn.microsoft.com/en-us/azure/defender-for-cloud/gain-end-user-context-ai
349-
// NOTE: The OpenAI .NET SDK (v2.7.0) does not currently expose a User property on ResponseCreationOptions.
350-
// When the SDK adds support, set: options.User = endUserId;
351-
_ = endUserId; // Suppress unused-variable warning until SDK support is available
351+
_ = endUserId;
352352

353353
// Add tools if provided
354354
if (tools != null)
@@ -503,6 +503,9 @@ private bool IsMcpToolAllowed(string toolName)
503503
[LoggerMessage(Level = LogLevel.Information, Message = "AI tool call invoked: tool={ToolName} iteration={Iteration}")]
504504
private static partial void LogMcpToolCallInvoked(ILogger logger, string toolName, int iteration);
505505

506+
[LoggerMessage(Level = LogLevel.Information, Message = "AI tool call invoked (streaming): tool={ToolName} depth={Depth}")]
507+
private static partial void LogMcpToolCallInvokedStream(ILogger logger, string toolName, int depth);
508+
506509
[LoggerMessage(Level = LogLevel.Warning, Message = "AI tool call rejected — not on allowlist: tool={ToolName}")]
507510
private static partial void LogMcpToolCallRejected(ILogger logger, string toolName);
508511

EssentialCSharp.Web.Tests/McpApiTokenServiceTests.cs

Lines changed: 85 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -71,4 +71,89 @@ public async Task CreateTokenAsync_WithExplicitCreatedAt_UsesReferenceTimeForDef
7171
await Assert.That(entity.ExpiresAt!.Value)
7272
.IsEqualTo(McpApiTokenService.GetDefaultExpirationUtc(createdAtUtc));
7373
}
74+
75+
[Test]
76+
public async Task GetActiveTokenCountAsync_NoTokens_ReturnsZero()
77+
{
78+
string userId = await McpTestHelper.CreateUserAsync(factory, "mcp-count-zero");
79+
80+
using var scope = factory.Services.CreateScope();
81+
var tokenService = scope.ServiceProvider.GetRequiredService<McpApiTokenService>();
82+
83+
int count = await tokenService.GetActiveTokenCountAsync(userId);
84+
85+
await Assert.That(count).IsEqualTo(0);
86+
}
87+
88+
[Test]
89+
public async Task GetActiveTokenCountAsync_ActiveTokens_CountsAll()
90+
{
91+
string userId = await McpTestHelper.CreateUserAsync(factory, "mcp-count-active");
92+
93+
using var scope = factory.Services.CreateScope();
94+
var tokenService = scope.ServiceProvider.GetRequiredService<McpApiTokenService>();
95+
96+
await tokenService.CreateTokenAsync(userId, "token-1");
97+
await tokenService.CreateTokenAsync(userId, "token-2");
98+
await tokenService.CreateTokenAsync(userId, "token-3");
99+
100+
int count = await tokenService.GetActiveTokenCountAsync(userId);
101+
102+
await Assert.That(count).IsEqualTo(3);
103+
}
104+
105+
[Test]
106+
public async Task GetActiveTokenCountAsync_RevokedToken_ExcludedFromCount()
107+
{
108+
string userId = await McpTestHelper.CreateUserAsync(factory, "mcp-count-revoked");
109+
110+
using var scope = factory.Services.CreateScope();
111+
var tokenService = scope.ServiceProvider.GetRequiredService<McpApiTokenService>();
112+
113+
await tokenService.CreateTokenAsync(userId, "active-token");
114+
(_, var revokedEntity) = await tokenService.CreateTokenAsync(userId, "revoked-token");
115+
await tokenService.RevokeTokenAsync(revokedEntity.Id, userId);
116+
117+
int count = await tokenService.GetActiveTokenCountAsync(userId);
118+
119+
await Assert.That(count).IsEqualTo(1);
120+
}
121+
122+
[Test]
123+
public async Task GetActiveTokenCountAsync_ExpiredToken_ExcludedFromCount()
124+
{
125+
string userId = await McpTestHelper.CreateUserAsync(factory, "mcp-count-expired");
126+
127+
using var scope = factory.Services.CreateScope();
128+
var tokenService = scope.ServiceProvider.GetRequiredService<McpApiTokenService>();
129+
130+
// Create a token that has already expired:
131+
// createdAt 7 months ago → max expiry = 1 month ago; use 2 months ago as expiresAt
132+
DateTime createdAt = DateTime.UtcNow.AddMonths(-7);
133+
DateTime pastExpiry = DateTime.UtcNow.AddMonths(-2);
134+
await tokenService.CreateTokenAsync(userId, "expired-token",
135+
expiresAt: pastExpiry, createdAtUtc: createdAt);
136+
await tokenService.CreateTokenAsync(userId, "valid-token");
137+
138+
int count = await tokenService.GetActiveTokenCountAsync(userId);
139+
140+
await Assert.That(count).IsEqualTo(1);
141+
}
142+
143+
[Test]
144+
public async Task CreateTokenAsync_AtMaxLimit_ThrowsInvalidOperationException()
145+
{
146+
string userId = await McpTestHelper.CreateUserAsync(factory, "mcp-at-limit");
147+
148+
using var scope = factory.Services.CreateScope();
149+
var tokenService = scope.ServiceProvider.GetRequiredService<McpApiTokenService>();
150+
151+
for (int i = 0; i < McpApiTokenService.MaxTokensPerUser; i++)
152+
{
153+
await tokenService.CreateTokenAsync(userId, $"token-{i}");
154+
}
155+
156+
await Assert.That(() => tokenService.CreateTokenAsync(userId, "one-too-many"))
157+
.Throws<InvalidOperationException>();
158+
}
74159
}
Lines changed: 123 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,123 @@
1+
using EssentialCSharp.Web.Services;
2+
using Microsoft.Extensions.Caching.Memory;
3+
4+
namespace EssentialCSharp.Web.Tests;
5+
6+
public class ResponseIdValidationServiceTests
7+
{
8+
private static ResponseIdValidationService CreateService()
9+
=> new(new MemoryCache(new MemoryCacheOptions()));
10+
11+
[Test]
12+
public async Task ValidateResponseId_NullResponseId_AllowsNewConversation()
13+
{
14+
var service = CreateService();
15+
16+
bool result = service.ValidateResponseId("user1", null);
17+
18+
await Assert.That(result).IsTrue();
19+
}
20+
21+
[Test]
22+
public async Task ValidateResponseId_EmptyResponseId_AllowsNewConversation()
23+
{
24+
var service = CreateService();
25+
26+
bool result = service.ValidateResponseId("user1", "");
27+
28+
await Assert.That(result).IsTrue();
29+
}
30+
31+
[Test]
32+
public async Task ValidateResponseId_NullUserId_Rejects()
33+
{
34+
var service = CreateService();
35+
36+
bool result = service.ValidateResponseId(null, "resp_123");
37+
38+
await Assert.That(result).IsFalse();
39+
}
40+
41+
[Test]
42+
public async Task ValidateResponseId_EmptyUserId_Rejects()
43+
{
44+
var service = CreateService();
45+
46+
bool result = service.ValidateResponseId("", "resp_123");
47+
48+
await Assert.That(result).IsFalse();
49+
}
50+
51+
[Test]
52+
public async Task ValidateResponseId_CacheMiss_AllowsGracefulDegradation()
53+
{
54+
var service = CreateService();
55+
// No RecordResponseId call — simulate server restart / different instance
56+
57+
bool result = service.ValidateResponseId("user1", "resp_unknown");
58+
59+
await Assert.That(result).IsTrue();
60+
}
61+
62+
[Test]
63+
public async Task ValidateResponseId_RecordedByOwner_Validates()
64+
{
65+
var service = CreateService();
66+
service.RecordResponseId("user1", "resp_abc");
67+
68+
bool result = service.ValidateResponseId("user1", "resp_abc");
69+
70+
await Assert.That(result).IsTrue();
71+
}
72+
73+
[Test]
74+
public async Task ValidateResponseId_RecordedByDifferentUser_Rejects()
75+
{
76+
var service = CreateService();
77+
service.RecordResponseId("user1", "resp_abc");
78+
79+
bool result = service.ValidateResponseId("user2", "resp_abc");
80+
81+
await Assert.That(result).IsFalse();
82+
}
83+
84+
[Test]
85+
public async Task RecordResponseId_NullInputs_DoesNotThrow()
86+
{
87+
var service = CreateService();
88+
89+
service.RecordResponseId(null, "resp_abc");
90+
service.RecordResponseId("user1", null);
91+
service.RecordResponseId(null, null);
92+
93+
// Verify the service is still functional after no-op calls
94+
service.RecordResponseId("user1", "resp_abc");
95+
await Assert.That(service.ValidateResponseId("user1", "resp_abc")).IsTrue();
96+
}
97+
98+
[Test]
99+
public async Task ValidateResponseId_MultipleResponseIds_EachValidatedIndependently()
100+
{
101+
var service = CreateService();
102+
service.RecordResponseId("user1", "resp_001");
103+
service.RecordResponseId("user1", "resp_002");
104+
105+
await Assert.That(service.ValidateResponseId("user1", "resp_001")).IsTrue();
106+
await Assert.That(service.ValidateResponseId("user1", "resp_002")).IsTrue();
107+
// Unrecorded ID for same user → cache miss → allow
108+
await Assert.That(service.ValidateResponseId("user1", "resp_003")).IsTrue();
109+
}
110+
111+
[Test]
112+
public async Task ValidateResponseId_TwoUsers_IsolatedFromEachOther()
113+
{
114+
var service = CreateService();
115+
service.RecordResponseId("user1", "resp_A");
116+
service.RecordResponseId("user2", "resp_B");
117+
118+
await Assert.That(service.ValidateResponseId("user1", "resp_A")).IsTrue();
119+
await Assert.That(service.ValidateResponseId("user2", "resp_B")).IsTrue();
120+
await Assert.That(service.ValidateResponseId("user2", "resp_A")).IsFalse();
121+
await Assert.That(service.ValidateResponseId("user1", "resp_B")).IsFalse();
122+
}
123+
}

EssentialCSharp.Web/Controllers/McpTokenController.cs

Lines changed: 21 additions & 18 deletions
Original file line numberDiff line numberDiff line change
@@ -22,10 +22,6 @@ public async Task<IActionResult> CreateToken(
2222
if (string.IsNullOrEmpty(userId))
2323
return Unauthorized(new { Error = "User must be logged in to generate an MCP token." });
2424

25-
int activeCount = await tokenService.GetActiveTokenCountAsync(userId, cancellationToken);
26-
if (activeCount >= McpApiTokenService.MaxTokensPerUser)
27-
return BadRequest(new { Error = $"You have reached the maximum of {McpApiTokenService.MaxTokensPerUser} active MCP tokens. Revoke an existing token before creating a new one." });
28-
2925
string name = string.IsNullOrWhiteSpace(request?.Name) ? "default" : request.Name.Trim();
3026
if (name.Length > 256)
3127
return BadRequest(new { Error = "Token name must be 256 characters or fewer." });
@@ -44,22 +40,29 @@ public async Task<IActionResult> CreateToken(
4440
expiresAt = expiresOn.ToDateTime(TimeOnly.MaxValue, DateTimeKind.Utc);
4541
}
4642

47-
var (rawToken, entity) = await tokenService.CreateTokenAsync(
48-
userId,
49-
name,
50-
expiresAt,
51-
createdAtUtc: nowUtc,
52-
cancellationToken: cancellationToken);
43+
try
44+
{
45+
var (rawToken, entity) = await tokenService.CreateTokenAsync(
46+
userId,
47+
name,
48+
expiresAt,
49+
createdAtUtc: nowUtc,
50+
cancellationToken: cancellationToken);
5351

54-
return Ok(new
52+
return Ok(new
53+
{
54+
TokenId = entity.Id,
55+
Token = rawToken,
56+
Name = entity.Name,
57+
ExpiresAt = entity.ExpiresAt,
58+
CreatedAt = entity.CreatedAt,
59+
Usage = "Add to your MCP client config: { \"url\": \"<site-url>/mcp\", \"headers\": { \"Authorization\": \"Bearer <token>\" } }"
60+
});
61+
}
62+
catch (InvalidOperationException ex)
5563
{
56-
TokenId = entity.Id,
57-
Token = rawToken,
58-
Name = entity.Name,
59-
ExpiresAt = entity.ExpiresAt,
60-
CreatedAt = entity.CreatedAt,
61-
Usage = "Add to your MCP client config: { \"url\": \"<site-url>/mcp\", \"headers\": { \"Authorization\": \"Bearer <token>\" } }"
62-
});
64+
return BadRequest(new { Error = ex.Message + $" Revoke an existing token before creating a new one." });
65+
}
6366
}
6467

6568
[HttpDelete("{id:guid}")]

EssentialCSharp.Web/Services/McpApiTokenService.cs

Lines changed: 14 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -32,6 +32,8 @@ public static string GenerateRawToken()
3232
/// <summary>
3333
/// Creates a new named API token for the specified user.
3434
/// Returns the raw token (shown once — never stored).
35+
/// Throws <see cref="InvalidOperationException"/> if the user is already at <see cref="MaxTokensPerUser"/>.
36+
/// The limit check and insert are wrapped in a serializable transaction to prevent races.
3537
/// </summary>
3638
public async Task<(string RawToken, McpApiToken Entity)> CreateTokenAsync(
3739
string userId,
@@ -40,6 +42,17 @@ public static string GenerateRawToken()
4042
DateTime? createdAtUtc = null,
4143
CancellationToken cancellationToken = default)
4244
{
45+
using var tx = await db.Database.BeginTransactionAsync(
46+
System.Data.IsolationLevel.Serializable, cancellationToken);
47+
48+
int activeCount = await GetActiveTokenCountAsync(userId, cancellationToken);
49+
if (activeCount >= MaxTokensPerUser)
50+
{
51+
await tx.RollbackAsync(cancellationToken);
52+
throw new InvalidOperationException(
53+
$"You have reached the maximum of {MaxTokensPerUser} active MCP tokens.");
54+
}
55+
4356
string raw = GenerateRawToken();
4457
DateTime createdAt = createdAtUtc ?? DateTime.UtcNow;
4558
DateTime effectiveExpiration = ResolveExpiration(expiresAt, createdAt);
@@ -54,6 +67,7 @@ public static string GenerateRawToken()
5467
};
5568
db.McpApiTokens.Add(entity);
5669
await db.SaveChangesAsync(cancellationToken);
70+
await tx.CommitAsync(cancellationToken);
5771
return (raw, entity);
5872
}
5973

0 commit comments

Comments
 (0)