using System.Data.Common; using System.Security.Cryptography; using System.Text; using Microsoft.EntityFrameworkCore.Diagnostics; namespace Jiaowu.Api.Infrastructure.Observability; public sealed class DatabaseCommandTelemetryInterceptor( ObservabilityOptions options, IHttpContextAccessor httpContextAccessor, ILogger logger) : DbCommandInterceptor { public override DbDataReader ReaderExecuted( DbCommand command, CommandExecutedEventData eventData, DbDataReader result) { Observe(command, eventData.Duration, "reader"); return result; } public override ValueTask ReaderExecutedAsync( DbCommand command, CommandExecutedEventData eventData, DbDataReader result, CancellationToken cancellationToken = default) { Observe(command, eventData.Duration, "reader"); return ValueTask.FromResult(result); } public override int NonQueryExecuted( DbCommand command, CommandExecutedEventData eventData, int result) { Observe(command, eventData.Duration, "nonquery"); return result; } public override ValueTask NonQueryExecutedAsync( DbCommand command, CommandExecutedEventData eventData, int result, CancellationToken cancellationToken = default) { Observe(command, eventData.Duration, "nonquery"); return ValueTask.FromResult(result); } public override object? ScalarExecuted( DbCommand command, CommandExecutedEventData eventData, object? result) { Observe(command, eventData.Duration, "scalar"); return result; } public override ValueTask ScalarExecutedAsync( DbCommand command, CommandExecutedEventData eventData, object? result, CancellationToken cancellationToken = default) { Observe(command, eventData.Duration, "scalar"); return ValueTask.FromResult(result); } public override void CommandFailed( DbCommand command, CommandErrorEventData eventData) => Observe( command, eventData.Duration, "failed", eventData.Exception.GetType().Name); public override Task CommandFailedAsync( DbCommand command, CommandErrorEventData eventData, CancellationToken cancellationToken = default) { Observe( command, eventData.Duration, "failed", eventData.Exception.GetType().Name); return Task.CompletedTask; } public override void CommandCanceled( DbCommand command, CommandEndEventData eventData) => Observe(command, eventData.Duration, "canceled", "canceled"); public override Task CommandCanceledAsync( DbCommand command, CommandEndEventData eventData, CancellationToken cancellationToken = default) { Observe(command, eventData.Duration, "canceled", "canceled"); return Task.CompletedTask; } private void Observe( DbCommand command, TimeSpan duration, string commandKind, string? errorType = null) { if (!options.Enabled) return; var queryName = GetQueryName(command.CommandText); var statementHash = GetStatementHash(command.CommandText); var provider = GetProviderName(command); var requestId = httpContextAccessor.HttpContext?.TraceIdentifier ?? "background"; var durationMilliseconds = duration.TotalMilliseconds; if (errorType is not null) { logger.LogError( "Database command failed after {DurationMs:F1} ms: " + "{QueryName} ({CommandKind}, {Provider}, hash {StatementHash}, " + "error {ErrorType}, request {RequestId}).", durationMilliseconds, queryName, commandKind, provider, statementHash, errorType, requestId); return; } if (durationMilliseconds < options.SlowQueryThresholdMilliseconds) return; if (options.IncludeSqlText) { logger.LogWarning( "Slow database command took {DurationMs:F1} ms: " + "{QueryName} ({CommandKind}, {Provider}, hash {StatementHash}, " + "request {RequestId}). " + "SQL template: {SqlTemplate}", durationMilliseconds, queryName, commandKind, provider, statementHash, requestId, Truncate(command.CommandText, options.MaximumSqlTextLength)); } else { logger.LogWarning( "Slow database command took {DurationMs:F1} ms: " + "{QueryName} ({CommandKind}, {Provider}, hash {StatementHash}, " + "request {RequestId}).", durationMilliseconds, queryName, commandKind, provider, statementHash, requestId); } } internal static string GetQueryName(string commandText) { using var reader = new StringReader(commandText); while (reader.ReadLine() is { } line) { var trimmed = line.Trim(); if (trimmed.Length == 0) continue; if (trimmed.StartsWith("-- ", StringComparison.Ordinal)) return Truncate(trimmed[3..].Trim(), 120); return $"{FirstToken(trimmed)}:{GetStatementHash(commandText)}"; } return $"unknown:{GetStatementHash(commandText)}"; } internal static string GetStatementHash(string commandText) { var bytes = SHA256.HashData(Encoding.UTF8.GetBytes(commandText)); return Convert.ToHexString(bytes.AsSpan(0, 6)).ToLowerInvariant(); } private static string FirstToken(string value) { var end = value.IndexOfAny([' ', '\t', '\r', '\n', '(']); var token = end < 0 ? value : value[..end]; return token.Length == 0 ? "command" : token.ToLowerInvariant(); } private static string GetProviderName(DbCommand command) { var typeName = command.GetType().FullName ?? command.GetType().Name; if (typeName.Contains("MySql", StringComparison.OrdinalIgnoreCase)) return "mysql"; if (typeName.Contains("Sqlite", StringComparison.OrdinalIgnoreCase)) return "sqlite"; return "other_sql"; } private static string Truncate(string value, int maximumLength) => value.Length <= maximumLength ? value : value[..maximumLength]; }