using System.Data.Common; using System.Diagnostics; using System.Diagnostics.Metrics; using System.Security.Cryptography; using System.Text; using Microsoft.EntityFrameworkCore.Diagnostics; namespace Jiaowu.Api.Infrastructure.Observability; public sealed class DatabaseCommandTelemetryInterceptor( ObservabilityOptions options, ILogger logger) : DbCommandInterceptor { public const string ActivitySourceName = "Jiaowu.Api.Database"; public const string MeterName = "Jiaowu.Api.Database"; private static readonly ActivitySource ActivitySource = new(ActivitySourceName); private static readonly Meter Meter = new(MeterName); private static readonly Histogram CommandDuration = Meter.CreateHistogram( "jiaowu.db.command.duration", "ms", "EF Core database command duration"); private static readonly Counter SlowCommandCount = Meter.CreateCounter( "jiaowu.db.command.slow", "{command}", "EF Core commands exceeding the configured slow-query threshold"); private static readonly Counter FailedCommandCount = Meter.CreateCounter( "jiaowu.db.command.failed", "{command}", "Failed EF Core database commands"); 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 traceId = Activity.Current?.TraceId.ToString() ?? "none"; var tags = new TagList { { "db.system.name", provider }, { "db.operation.name", commandKind }, { "db.query.name", queryName } }; if (errorType is not null) tags.Add("error.type", errorType); var durationMilliseconds = duration.TotalMilliseconds; CommandDuration.Record(durationMilliseconds, tags); if (errorType is not null) FailedCommandCount.Add(1, tags); using var activity = ActivitySource.StartActivity( ActivityKind.Client, Activity.Current?.Context ?? default, startTime: DateTimeOffset.UtcNow - duration, name: queryName); if (activity is not null) { activity.SetTag("db.system.name", provider); activity.SetTag("db.operation.name", commandKind); activity.SetTag("db.query.name", queryName); activity.SetTag("db.statement.hash", statementHash); activity.SetTag( "db.namespace", EmptyToNull(command.Connection?.Database)); if (options.IncludeSqlText) { activity.SetTag( "db.query.text", Truncate(command.CommandText, options.MaximumSqlTextLength)); } if (errorType is not null) { activity.SetTag("error.type", errorType); activity.SetStatus(ActivityStatusCode.Error, errorType); } activity.SetEndTime(DateTime.UtcNow); } if (errorType is not null) { logger.LogError( "Database command failed after {DurationMs:F1} ms: " + "{QueryName} ({CommandKind}, {Provider}, hash {StatementHash}, " + "error {ErrorType}, trace {TraceId}).", durationMilliseconds, queryName, commandKind, provider, statementHash, errorType, traceId); return; } if (durationMilliseconds < options.SlowQueryThresholdMilliseconds) return; SlowCommandCount.Add(1, tags); if (options.IncludeSqlText) { logger.LogWarning( "Slow database command took {DurationMs:F1} ms: " + "{QueryName} ({CommandKind}, {Provider}, hash {StatementHash}, " + "trace {TraceId}). " + "SQL template: {SqlTemplate}", durationMilliseconds, queryName, commandKind, provider, statementHash, traceId, Truncate(command.CommandText, options.MaximumSqlTextLength)); } else { logger.LogWarning( "Slow database command took {DurationMs:F1} ms: " + "{QueryName} ({CommandKind}, {Provider}, hash {StatementHash}, " + "trace {TraceId}).", durationMilliseconds, queryName, commandKind, provider, statementHash, traceId); } } 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? EmptyToNull(string? value) => string.IsNullOrWhiteSpace(value) ? null : value; private static string Truncate(string value, int maximumLength) => value.Length <= maximumLength ? value : value[..maximumLength]; }