using System.Linq.Expressions; using System.Text; namespace Strata.SqlTools.Visitors.LinqToSql; /// /// Expression visitor for analyzing LINQ to SQL expression trees. /// Extracts query components such as SELECT, WHERE, JOIN, GROUP BY, and ORDER BY. /// public class LinqExpressionVisitor : ExpressionVisitor { private readonly StringBuilder _whereBuilder = new(); private readonly StringBuilder _orderByBuilder = new(); private readonly List _methodCalls = []; private bool _isInWhereClause; /// /// Gets the SELECT clause extracted from the expression. /// public string? SelectClause { get; private set; } /// /// Gets the FROM clause (table name) extracted from the expression. /// public string? FromClause { get; private set; } /// /// Gets the WHERE clause extracted from the expression. /// public string? WhereClause { get; private set; } /// /// Gets the ORDER BY clause extracted from the expression. /// public string? OrderByClause { get; private set; } /// /// Gets the GROUP BY clause extracted from the expression. /// public string? GroupByClause { get; private set; } /// /// Gets the list of LINQ method calls in the query chain. /// public List MethodCallChain => _methodCalls; /// /// Visits a method call expression. /// protected override Expression VisitMethodCall(MethodCallExpression node) { var methodName = node.Method.Name; _methodCalls.Add(methodName); switch (methodName) { case "Where": VisitWhereMethod(node); break; case "Select": VisitSelectMethod(node); break; case "OrderBy": case "OrderByDescending": case "ThenBy": case "ThenByDescending": VisitOrderByMethod(node); break; case "GroupBy": VisitGroupByMethod(node); break; case "Join": case "GroupJoin": VisitJoinMethod(node); break; case "Take": case "Skip": VisitTakeSkipMethod(node); break; default: // Visit the source expression Visit(node.Arguments[0]); break; } return node; } /// /// Visits a constant expression to extract the table name. /// protected override Expression VisitConstant(ConstantExpression node) { // Handle WHERE clause constants if (_isInWhereClause) { if (node.Value is string) { _whereBuilder.Append($"'{node.Value}'"); } else if (node.Value != null) { _whereBuilder.Append(node.Value.ToString()); } else { _whereBuilder.Append("NULL"); } return node; } // Handle table name extraction if (node.Type.IsGenericType) { var genericType = node.Type.GetGenericTypeDefinition(); if (genericType.Name.Contains("Table") || genericType.Name.Contains("Query")) { var entityType = node.Type.GetGenericArguments().FirstOrDefault(); if (entityType != null) { FromClause = entityType.Name; } } } return base.VisitConstant(node); } private void VisitWhereMethod(MethodCallExpression node) { // Visit the source Visit(node.Arguments[0]); // Extract the predicate if (node.Arguments.Count > 1) { var lambda = StripQuotes(node.Arguments[1]) as LambdaExpression; if (lambda != null) { _isInWhereClause = true; Visit(lambda.Body); _isInWhereClause = false; if (_whereBuilder.Length > 0) { WhereClause = _whereBuilder.ToString(); } } } } private void VisitSelectMethod(MethodCallExpression node) { // Visit the source Visit(node.Arguments[0]); // Extract the selector if (node.Arguments.Count > 1) { var lambda = StripQuotes(node.Arguments[1]) as LambdaExpression; if (lambda != null) { var selectExpression = ExtractSelectExpression(lambda.Body); if (!string.IsNullOrEmpty(selectExpression)) { SelectClause = selectExpression; } } } } private void VisitOrderByMethod(MethodCallExpression node) { // Visit the source Visit(node.Arguments[0]); // Extract the key selector if (node.Arguments.Count > 1) { var lambda = StripQuotes(node.Arguments[1]) as LambdaExpression; if (lambda != null) { var orderByExpression = ExtractMemberName(lambda.Body); if (!string.IsNullOrEmpty(orderByExpression)) { var direction = node.Method.Name.Contains("Descending") ? " DESC" : " ASC"; if (_orderByBuilder.Length > 0) { _orderByBuilder.Append(", "); } _orderByBuilder.Append(orderByExpression + direction); OrderByClause = _orderByBuilder.ToString(); } } } } private void VisitGroupByMethod(MethodCallExpression node) { // Visit the source Visit(node.Arguments[0]); // Extract the key selector if (node.Arguments.Count > 1) { var lambda = StripQuotes(node.Arguments[1]) as LambdaExpression; if (lambda != null) { var groupByExpression = ExtractMemberName(lambda.Body); if (!string.IsNullOrEmpty(groupByExpression)) { GroupByClause = groupByExpression; } } } } private void VisitJoinMethod(MethodCallExpression node) { // Visit the source Visit(node.Arguments[0]); // For joins, we'd need more complex logic to extract full join information // This is a simplified version _methodCalls.Add($"{node.Method.Name} (complex join analysis not fully implemented)"); } private void VisitTakeSkipMethod(MethodCallExpression node) { // Visit the source Visit(node.Arguments[0]); // Extract the count if (node.Arguments.Count > 1 && node.Arguments[1] is ConstantExpression constant) { _methodCalls.Add($"{node.Method.Name}({constant.Value})"); } } /// /// Visits a binary expression (e.g., comparisons, logical operations). /// protected override Expression VisitBinary(BinaryExpression node) { if (_isInWhereClause) { _whereBuilder.Append('('); Visit(node.Left); _whereBuilder.Append($" {GetOperator(node.NodeType)} "); Visit(node.Right); _whereBuilder.Append(')'); return node; } return base.VisitBinary(node); } /// /// Visits a member access expression. /// protected override Expression VisitMember(MemberExpression node) { if (_isInWhereClause) { var memberName = GetFullMemberName(node); _whereBuilder.Append(memberName); return node; } return base.VisitMember(node); } private static string ExtractSelectExpression(Expression expression) { if (expression is NewExpression newExpr) { var members = new List(); for (int i = 0; i < newExpr.Arguments.Count; i++) { var memberName = ExtractMemberName(newExpr.Arguments[i]); var alias = newExpr.Members?[i].Name; if (!string.IsNullOrEmpty(alias) && alias != memberName) { members.Add($"{memberName} AS {alias}"); } else { members.Add(memberName); } } return string.Join(", ", members); } var name = ExtractMemberName(expression); return string.IsNullOrEmpty(name) ? "*" : name; } private static string ExtractMemberName(Expression expression) { if (expression is MemberExpression member) { return GetFullMemberName(member); } if (expression is ParameterExpression) { return "*"; } if (expression is MethodCallExpression methodCall) { return $"{methodCall.Method.Name}(...)"; } return expression.ToString(); } private static string GetFullMemberName(MemberExpression expression) { var parts = new Stack(); var current = expression; while (current != null) { parts.Push(current.Member.Name); if (current.Expression is MemberExpression memberExpr) { current = memberExpr; } else if (current.Expression is ParameterExpression paramExpr) { // Use parameter name as table alias if it's not the default if (paramExpr.Name != null && paramExpr.Name.Length == 1) { parts.Push(paramExpr.Name); } break; } else { break; } } return string.Join(".", parts); } private static string GetOperator(ExpressionType nodeType) { return nodeType switch { ExpressionType.Equal => "=", ExpressionType.NotEqual => "!=", ExpressionType.GreaterThan => ">", ExpressionType.GreaterThanOrEqual => ">=", ExpressionType.LessThan => "<", ExpressionType.LessThanOrEqual => "<=", ExpressionType.AndAlso => "AND", ExpressionType.OrElse => "OR", ExpressionType.Add => "+", ExpressionType.Subtract => "-", ExpressionType.Multiply => "*", ExpressionType.Divide => "/", _ => nodeType.ToString() }; } private static Expression StripQuotes(Expression expression) { while (expression.NodeType == ExpressionType.Quote) { expression = ((UnaryExpression)expression).Operand; } return expression; } }