diff --git a/.env.example b/.env.example index be8ea9e..04e74d8 100644 --- a/.env.example +++ b/.env.example @@ -141,6 +141,9 @@ EMBEDDING_FAIL_OPEN=false # AGENT_WEBUI_URL Agent HTTP API. Compose default: http://deepsql-agent:8787 # AGENT_PROVISIONER_URL Per-user profile provisioner. Compose default: # http://deepsql-agent:8788/provision +# DEEPSQL_API_BASE_URL Where the agent container's MCP tools call the backend. +# Compose default: http://backend:8080/api/ +# Native Java + Compose agent: http://host.docker.internal:8080/api/ # AGENT_PROVISION_SECRET Shared secret between backend and agent (required). # DEEPSQL_AGENT_PORT / DEEPSQL_AGENT_PROVISIONER_PORT — host port mappings # (compose binds these to 127.0.0.1 only; public path is nginx /agent-api). diff --git a/AGENTS.md b/AGENTS.md index 072bf15..cff7764 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -186,6 +186,11 @@ only covers cloud-specific, non-obvious caveats. `HERMES_WEBUI_ALLOWED_ORIGINS` (upstream env name). - A demo target DB `demo_shop` (same Postgres server, sample `customers`/`products`/`orders`) exists for exercising connection/schema features without an external database. +- A multi-schema fixture DB `acme_erp` (schemas: `crm`, `sales`, `finance`, `inventory`, + `hr`, `marts`) exists for chat-access-policy and multi-schema tests. Seed with: + `sudo -u postgres psql -f docker/postgres/init/11_create_acme_erp.sql` then + `bash scripts/seed-acme-erp.sh` (registers `ACME ERP (Multi-Schema)` when backend auth + is disabled or you have an admin session cookie). ### Non-obvious setup caveats (each cost real debugging time) @@ -259,9 +264,16 @@ only covers cloud-specific, non-obvious caveats. non-`public` schemas (`crm`, `sales`, `finance`, `hr`, `inventory`) for Brain / MCP cross-schema checks. Prefer schema-qualified SQL (`sales.orders`); bare names follow the role’s `search_path` (usually `public`). -- **`AGENT_WEBUI_URL` for native runs.** Default is `http://deepsql-agent:8787` - (Compose DNS). Native local must set `AGENT_WEBUI_URL=http://127.0.0.1:8787` in - `.env` or CLI/Slack `AgentChatClient` cannot reach the agent API. +- **`AGENT_WEBUI_URL` / `AGENT_PROVISIONER_URL` for native runs.** Compose + defaults (`http://deepsql-agent:8787` and `…:8788/provision`) do not resolve + on the host. Native local must point both at loopback + (`http://127.0.0.1:8787` and `http://127.0.0.1:8788/provision`) or the Agent + tab returns 503 `Could not provision the DeepSQL Agent for this user`. + `scripts/start-backend.sh` remaps those hostnames automatically when they + don't resolve. If the agent container is used with a host-side Java backend, + set `DEEPSQL_API_BASE_URL=http://host.docker.internal:8080/api/` so MCP + tools can reach the native process (compose publishes `host.docker.internal` + via `extra_hosts`). - **DeepSQL CLI (`deepsql`) for agent testing.** Install from the repo package: `cd mcp && DEEPSQL_SKIP_AGENT_SETUP=1 npm install -g .` (prefix `~/.npm-global`, keep that on `PATH`). Auth against local backend with an MCP diff --git a/CLAUDE.md b/CLAUDE.md index 2cf64a0..06cd7ce 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -93,7 +93,7 @@ backend/ repository/ # Spring Data repositories provider/ # Database dialect registry (PostgreSQL, MySQL) config/ # Spring configuration - security/ # JWT auth, RBAC + security/ # JWT auth, RBAC, admin profile switch (`ImpersonationService`) llm/ # LLM provider registry, config resolver, OpenAI-compatible provider util/ # Shared utilities src/test/ # JUnit 5 tests @@ -110,7 +110,7 @@ src/ # Frontend (React) components/ # UI components tabs/ # 40+ specialized tabs sections/ # Top-level sidebar destinations (Agent, Dashboards, Brain, - # Performance = Slow Queries + Workload, Editor, Docs) + # Performance = Slow Queries + Workload, Editor) lib/ api/client.js # Centralized API layer (axios, 25+ modules) stores/ # Zustand stores (dashboard, connection, chat, UI) @@ -209,6 +209,9 @@ returns a number). 4. **Tooltips**: Always use `HelpTooltip` component, never plain `title` attributes. 5. **Design**: Minimal black/white/grey palette, Inter font, subtle transitions. See UX guidelines in full CLAUDE.md. +### Admin profile switch +Admins can **View as** a sub-user from the top-right of the home layout (`ProfileSwitch`) to verify connection ACLs, chat/editor policies, and role-gated nav. The admin JWT stays on the session; `ImpersonationService` sets an httpOnly `impersonate_user` cookie and `JwtAuthenticationFilter` overlays the target principal. `POST|DELETE|GET /api/admin/impersonate` are excluded from the overlay so stop/list still run as the real admin. Cannot target another ADMIN, self, or a non-ACTIVE account. `/auth/me` returns the **effective** user plus `impersonating` / `impersonatorUsername`. + ### Git Rules - Do NOT commit automatically — wait for explicit user instruction. - Conventional commits: `feat:`, `fix:`, `refactor:`, `docs:`, `test:`, `chore:`, `perf:`, `ci:` diff --git a/backend/src/main/java/com/dbaagent/controller/AuthController.java b/backend/src/main/java/com/dbaagent/controller/AuthController.java index 84a5b4b..e5d3cf7 100644 --- a/backend/src/main/java/com/dbaagent/controller/AuthController.java +++ b/backend/src/main/java/com/dbaagent/controller/AuthController.java @@ -9,6 +9,8 @@ import com.dbaagent.service.AuthSessionService; import com.dbaagent.service.PasswordlessAuthService; import com.dbaagent.service.PermissionService; +import com.dbaagent.service.ImpersonationService; +import com.dbaagent.security.ImpersonationContext; import com.dbaagent.service.SystemConfigService; import com.dbaagent.service.UserInviteService; import jakarta.servlet.http.Cookie; @@ -47,6 +49,7 @@ public class AuthController { private final PrivateBetaRequestRepository privateBetaRequestRepository; private final SystemConfigService systemConfigService; private final AgentBridgeService agentBridgeService; + private final ImpersonationService impersonationService; @Value("${security.cookie.refresh-name:refresh_token}") private String refreshCookieName; @@ -201,12 +204,21 @@ public ResponseEntity refreshSession(HttpServletRequest httpRequest, HttpServ return ResponseEntity.status(401).body(Map.of("message", "Session expired")); } authSessionService.writeSessionCookies(httpResponse, refreshed.get()); + User effectiveUser = impersonationService.resolveFromCookie(httpRequest, user) + .map(ImpersonationContext.State::target) + .orElse(user); // Keep the user's agent token alive for as long as the UI session lives. // The SPA refreshes on access-token expiry (~every 15 min of activity), so // this slides the agent token forward on each active interval — a logged-in // UI never ends up with a dead agent. agentBridgeService.extendAgentTokens(user.getUsername()); - return ResponseEntity.ok(toAuthPayload(user, user.getRoleEnum(), permissionService.getEffectivePermissionCodes(user.getRoleEnum()))); + Map payload = toAuthPayload( + effectiveUser, + effectiveUser.getRoleEnum(), + permissionService.getEffectivePermissionCodes(effectiveUser.getRoleEnum()) + ); + impersonationService.decorateAuthPayload(httpRequest, user, payload); + return ResponseEntity.ok(payload); } @PostMapping("/logout") @@ -304,7 +316,7 @@ public ResponseEntity acceptInvite( } @GetMapping("/me") - public ResponseEntity getCurrentUser() { + public ResponseEntity getCurrentUser(HttpServletRequest httpRequest) { Authentication auth = SecurityContextHolder.getContext().getAuthentication(); if (auth == null || !auth.isAuthenticated() || "anonymousUser".equals(auth.getPrincipal())) { return ResponseEntity.status(401).body(Map.of("message", "Not authenticated")); @@ -313,9 +325,7 @@ public ResponseEntity getCurrentUser() { Role role = user.getRoleEnum(); Set permissions = permissionService.getEffectivePermissionCodes(role); Map response = toAuthPayload(user, role, permissions); - response.put("emailVerified", user.isEmailVerified()); - response.put("accountStatus", user.getAccountStatus()); - response.put("emailTwoFactorEnabled", systemConfigService.getBoolean("security.workspace.email2fa.enabled")); + impersonationService.decorateAuthPayload(httpRequest, user, response); return ResponseEntity.ok(response); } @@ -360,6 +370,7 @@ private ResponseEntity authResponse(PasswordlessAuthService.AuthFlowResult re } if (result.sessionAuthentication() != null && result.user() != null && result.role() != null) { authSessionService.writeSessionCookies(httpResponse, result.sessionAuthentication()); + authSessionService.clearImpersonationCookie(httpResponse); Set permissionNames = result.permissions() == null ? Set.of() : result.permissions().stream() .map(Enum::name) .collect(Collectors.toSet()); diff --git a/backend/src/main/java/com/dbaagent/controller/ImpersonationController.java b/backend/src/main/java/com/dbaagent/controller/ImpersonationController.java new file mode 100644 index 0000000..8810972 --- /dev/null +++ b/backend/src/main/java/com/dbaagent/controller/ImpersonationController.java @@ -0,0 +1,136 @@ +package com.dbaagent.controller; + +import com.dbaagent.model.Role; +import com.dbaagent.model.User; +import com.dbaagent.repository.UserRepository; +import com.dbaagent.security.ImpersonationContext; +import com.dbaagent.service.ImpersonationService; +import com.dbaagent.service.PermissionService; +import jakarta.servlet.http.HttpServletRequest; +import jakarta.servlet.http.HttpServletResponse; +import lombok.RequiredArgsConstructor; +import org.springframework.http.ResponseEntity; +import org.springframework.security.access.prepost.PreAuthorize; +import org.springframework.security.core.Authentication; +import org.springframework.security.core.context.SecurityContextHolder; +import org.springframework.web.bind.annotation.DeleteMapping; +import org.springframework.web.bind.annotation.GetMapping; +import org.springframework.web.bind.annotation.PostMapping; +import org.springframework.web.bind.annotation.RequestBody; +import org.springframework.web.bind.annotation.RequestMapping; +import org.springframework.web.bind.annotation.RestController; +import org.springframework.web.server.ResponseStatusException; + +import java.util.LinkedHashMap; +import java.util.Map; +import java.util.Set; + +/** + * Admin-only profile switch. These paths are excluded from the impersonation + * overlay so the caller stays the real administrator while starting, listing, + * or stopping a switch. + */ +@RestController +@RequestMapping("/admin/impersonate") +@PreAuthorize("hasRole('ADMIN')") +@RequiredArgsConstructor +public class ImpersonationController { + + private final ImpersonationService impersonationService; + private final UserRepository userRepository; + private final PermissionService permissionService; + + @GetMapping + public ResponseEntity> status(HttpServletRequest request) { + User actor = currentAdmin(); + ImpersonationContext.State state = impersonationService.resolveFromCookie(request, actor).orElse(null); + Map body = new LinkedHashMap<>(); + body.put("impersonating", state != null); + body.put("impersonator", Map.of( + "id", actor.getId(), + "username", actor.getUsername(), + "email", actor.getEmail() + )); + body.put("target", state == null ? null : candidateView(state.target())); + body.put("candidates", impersonationService.listCandidates(actor)); + return ResponseEntity.ok(body); + } + + @PostMapping + public ResponseEntity> start( + @RequestBody Map requestBody, + HttpServletRequest request, + HttpServletResponse response + ) { + User actor = currentAdmin(); + Long userId = readUserId(requestBody); + ImpersonationContext.State state = impersonationService.start(actor, userId, request, response); + return ResponseEntity.ok(toAuthPayload(state.target(), actor)); + } + + @DeleteMapping + public ResponseEntity> stop( + HttpServletRequest request, + HttpServletResponse response + ) { + User actor = currentAdmin(); + User restored = impersonationService.stop(actor, request, response); + Map payload = toAuthPayload(restored, null); + payload.put("impersonating", false); + return ResponseEntity.ok(payload); + } + + private Map toAuthPayload(User user, User impersonator) { + Role role = user.getRoleEnum(); + Set permissions = permissionService.getEffectivePermissionCodes(role); + Map payload = new LinkedHashMap<>(); + payload.put("username", user.getUsername()); + payload.put("email", user.getEmail()); + payload.put("role", role.name()); + payload.put("permissions", permissions); + payload.put("emailVerified", user.isEmailVerified()); + payload.put("accountStatus", user.getAccountStatus()); + if (impersonator != null) { + payload.put("impersonating", true); + payload.put("impersonatorUsername", impersonator.getUsername()); + payload.put("impersonatorEmail", impersonator.getEmail()); + } else { + payload.put("impersonating", false); + } + return payload; + } + + private Map candidateView(User user) { + Map dto = new LinkedHashMap<>(); + dto.put("id", user.getId()); + dto.put("username", user.getUsername()); + dto.put("email", user.getEmail()); + dto.put("role", user.getRole()); + dto.put("accountStatus", user.getAccountStatus()); + return dto; + } + + private Long readUserId(Map requestBody) { + if (requestBody == null || requestBody.get("userId") == null) { + return null; + } + Object raw = requestBody.get("userId"); + if (raw instanceof Number number) { + return number.longValue(); + } + try { + return Long.parseLong(String.valueOf(raw).trim()); + } catch (NumberFormatException e) { + return null; + } + } + + private User currentAdmin() { + Authentication auth = SecurityContextHolder.getContext().getAuthentication(); + if (auth == null || !auth.isAuthenticated() || "anonymousUser".equals(auth.getPrincipal())) { + throw new ResponseStatusException(org.springframework.http.HttpStatus.UNAUTHORIZED, "Not authenticated"); + } + return userRepository.findByUsername(auth.getName()) + .orElseThrow(() -> new ResponseStatusException(org.springframework.http.HttpStatus.UNAUTHORIZED, "User not found")); + } +} diff --git a/backend/src/main/java/com/dbaagent/controller/SchemaController.java b/backend/src/main/java/com/dbaagent/controller/SchemaController.java index 72204be..ca3064f 100644 --- a/backend/src/main/java/com/dbaagent/controller/SchemaController.java +++ b/backend/src/main/java/com/dbaagent/controller/SchemaController.java @@ -11,6 +11,7 @@ import com.dbaagent.service.RunningQueryRegistry; import com.dbaagent.service.SqlExecutionAuditService; import com.dbaagent.service.UserDataAccessPolicyException; +import com.dbaagent.service.UserDataAccessPolicyService; import com.dbaagent.service.SchemaScannerService; import com.dbaagent.service.VisualizationService; import com.dbaagent.service.security.AccessControlService; @@ -38,6 +39,7 @@ public class SchemaController { private final QueryExecutorService queryExecutorService; private final AccessControlService accessControlService; private final SqlExecutionAuditService sqlExecutionAuditService; + private final UserDataAccessPolicyService userDataAccessPolicyService; private final RunningQueryRegistry runningQueryRegistry; private final ActiveQueryService activeQueryService; @@ -52,7 +54,7 @@ public ResponseEntity> scanSchema(@PathVariable String conne } accessControlService.assertCanUseChatEditor(connectionId); - SchemaMetadata schema = schemaScannerService.scanSchema(connectionId); + SchemaMetadata schema = scopedSchema(connectionId, schemaScannerService.scanSchema(connectionId)); response.put("success", true); response.put("schema", schema); return ResponseEntity.ok(response); @@ -82,7 +84,7 @@ public ResponseEntity> getSchema(@PathVariable String connec } accessControlService.assertCanUseChatEditor(connectionId); - SchemaMetadata schema = schemaScannerService.scanSchema(connectionId); + SchemaMetadata schema = scopedSchema(connectionId, schemaScannerService.scanSchema(connectionId)); response.put("success", true); response.put("schema", schema); return ResponseEntity.ok(response); @@ -112,7 +114,7 @@ public ResponseEntity> getVisualization(@PathVariable String } accessControlService.assertCanUseChatEditor(connectionId); - SchemaMetadata schema = schemaScannerService.scanSchema(connectionId); + SchemaMetadata schema = scopedSchema(connectionId, schemaScannerService.scanSchema(connectionId)); ErDiagramData erDiagram = visualizationService.generateErDiagram(schema, connectionId); DependencyGraphData dependencyGraph = visualizationService.generateDependencyGraph(schema, connectionId); @@ -147,7 +149,10 @@ public ResponseEntity> getDatabaseObjects(@PathVariable Stri } accessControlService.assertCanUseChatEditor(connectionId); - List objects = queryExecutorService.getDatabaseObjects(connectionId); + List objects = scopedObjects( + connectionId, + queryExecutorService.getDatabaseObjects(connectionId) + ); log.info("Successfully fetched {} database objects for connection: {}", objects.size(), connectionId); response.put("success", true); response.put("objects", objects); @@ -334,6 +339,12 @@ public ResponseEntity> getTableIndexes( return ResponseEntity.status(HttpStatus.NOT_FOUND).body(response); } accessControlService.assertCanUseChatEditor(connectionId); + userDataAccessPolicyService.assertTableSchemaAllowed( + connectionId, + accessControlService.getCurrentUsername(), + accessControlService.isCurrentUserAdmin(), + tableName + ); List indexes = queryExecutorService.getTableIndexes(connectionId, tableName); response.put("success", true); @@ -343,6 +354,11 @@ public ResponseEntity> getTableIndexes( response.put("success", false); response.put("message", "Failed to fetch table indexes: " + e.getMessage()); return ResponseEntity.status(HttpStatus.INTERNAL_SERVER_ERROR).body(response); + } catch (UserDataAccessPolicyException e) { + response.put("success", false); + response.put("message", e.getMessage()); + response.put("errorCode", e.getErrorCode()); + return ResponseEntity.status(HttpStatus.FORBIDDEN).body(response); } catch (ResponseStatusException e) { response.put("success", false); response.put("message", e.getReason()); @@ -366,6 +382,12 @@ public ResponseEntity> getTableStats( return ResponseEntity.status(HttpStatus.NOT_FOUND).body(response); } accessControlService.assertCanUseChatEditor(connectionId); + userDataAccessPolicyService.assertTableSchemaAllowed( + connectionId, + accessControlService.getCurrentUsername(), + accessControlService.isCurrentUserAdmin(), + tableName + ); TableStats stats = queryExecutorService.getTableStats(connectionId, tableName); response.put("success", true); @@ -375,6 +397,11 @@ public ResponseEntity> getTableStats( response.put("success", false); response.put("message", "Failed to fetch table stats: " + e.getMessage()); return ResponseEntity.status(HttpStatus.INTERNAL_SERVER_ERROR).body(response); + } catch (UserDataAccessPolicyException e) { + response.put("success", false); + response.put("message", e.getMessage()); + response.put("errorCode", e.getErrorCode()); + return ResponseEntity.status(HttpStatus.FORBIDDEN).body(response); } catch (ResponseStatusException e) { response.put("success", false); response.put("message", e.getReason()); @@ -385,4 +412,22 @@ public ResponseEntity> getTableStats( return ResponseEntity.status(HttpStatus.INTERNAL_SERVER_ERROR).body(response); } } + + private SchemaMetadata scopedSchema(String connectionId, SchemaMetadata schema) { + return userDataAccessPolicyService.filterSchemaMetadata( + connectionId, + accessControlService.getCurrentUsername(), + accessControlService.isCurrentUserAdmin(), + schema + ); + } + + private List scopedObjects(String connectionId, List objects) { + return userDataAccessPolicyService.filterDatabaseObjects( + connectionId, + accessControlService.getCurrentUsername(), + accessControlService.isCurrentUserAdmin(), + objects + ); + } } diff --git a/backend/src/main/java/com/dbaagent/model/SecurityEventType.java b/backend/src/main/java/com/dbaagent/model/SecurityEventType.java index 8ad7962..7559ee2 100644 --- a/backend/src/main/java/com/dbaagent/model/SecurityEventType.java +++ b/backend/src/main/java/com/dbaagent/model/SecurityEventType.java @@ -25,6 +25,8 @@ public enum SecurityEventType { SESSION_REFRESHED, SESSION_REVOKED, SESSION_EXPIRED, + IMPERSONATION_STARTED, + IMPERSONATION_STOPPED, LOGOUT, LOGOUT_ALL, ACCOUNT_LOCKED, diff --git a/backend/src/main/java/com/dbaagent/security/ImpersonationContext.java b/backend/src/main/java/com/dbaagent/security/ImpersonationContext.java new file mode 100644 index 0000000..240354f --- /dev/null +++ b/backend/src/main/java/com/dbaagent/security/ImpersonationContext.java @@ -0,0 +1,50 @@ +package com.dbaagent.security; + +import com.dbaagent.model.User; + +import java.util.Optional; + +/** + * Request-scoped impersonation overlay. The admin JWT stays on the session; + * {@link JwtAuthenticationFilter} swaps the SecurityContext principal to the + * target user and records both identities here so {@code /auth/me} can show a + * banner and {@code AccessControlService} can honour the target even when + * {@code security.auth.enabled} is false. + */ +public final class ImpersonationContext { + + public record State(User impersonator, User target) { + public String impersonatorUsername() { + return impersonator != null ? impersonator.getUsername() : null; + } + + public String impersonatorEmail() { + return impersonator != null ? impersonator.getEmail() : null; + } + + public String targetUsername() { + return target != null ? target.getUsername() : null; + } + } + + private static final ThreadLocal CURRENT = new ThreadLocal<>(); + + private ImpersonationContext() { + } + + public static void enter(State state) { + CURRENT.set(state); + } + + public static void clear() { + CURRENT.remove(); + } + + public static Optional current() { + return Optional.ofNullable(CURRENT.get()); + } + + public static boolean isActive() { + return CURRENT.get() != null; + } +} diff --git a/backend/src/main/java/com/dbaagent/security/JwtAuthenticationFilter.java b/backend/src/main/java/com/dbaagent/security/JwtAuthenticationFilter.java index ddecb90..9262c33 100644 --- a/backend/src/main/java/com/dbaagent/security/JwtAuthenticationFilter.java +++ b/backend/src/main/java/com/dbaagent/security/JwtAuthenticationFilter.java @@ -16,6 +16,7 @@ import org.springframework.web.filter.OncePerRequestFilter; import com.dbaagent.service.AuthSessionService; +import com.dbaagent.service.ImpersonationService; import java.io.IOException; import java.util.ArrayList; @@ -40,6 +41,9 @@ public class JwtAuthenticationFilter extends OncePerRequestFilter { @Autowired private AuthSessionService authSessionService; + @Autowired + private ImpersonationService impersonationService; + @Override protected void doFilterInternal(HttpServletRequest request, HttpServletResponse response, FilterChain chain) throws ServletException, IOException { @@ -69,7 +73,7 @@ protected void doFilterInternal(HttpServletRequest request, HttpServletResponse UsernamePasswordAuthenticationToken auth = new UsernamePasswordAuthenticationToken( "admin", null, devAuthorities); SecurityContextHolder.getContext().setAuthentication(auth); - chain.doFilter(request, response); + applyImpersonation(request, response, chain); return; } @@ -132,7 +136,17 @@ protected void doFilterInternal(HttpServletRequest request, HttpServletResponse } else if (isBrainRequest) { log.warn("Brain auth missing user: path={}, username={}", requestPath, username); } - chain.doFilter(request, response); + applyImpersonation(request, response, chain); + } + + private void applyImpersonation(HttpServletRequest request, HttpServletResponse response, FilterChain chain) + throws ServletException, IOException { + try { + impersonationService.applyToRequest(request); + chain.doFilter(request, response); + } finally { + ImpersonationContext.clear(); + } } private String extractUsernameSafely(String token) { diff --git a/backend/src/main/java/com/dbaagent/service/AuthSessionService.java b/backend/src/main/java/com/dbaagent/service/AuthSessionService.java index 5f85e30..1e40548 100644 --- a/backend/src/main/java/com/dbaagent/service/AuthSessionService.java +++ b/backend/src/main/java/com/dbaagent/service/AuthSessionService.java @@ -38,6 +38,9 @@ public class AuthSessionService { @Value("${security.cookie.refresh-name:refresh_token}") private String refreshCookieName; + @Value("${security.cookie.impersonate-name:impersonate_user}") + private String impersonateCookieName; + @Value("${security.cookie.secure:false}") private boolean cookieSecure; @@ -152,9 +155,29 @@ public void writeSessionCookies(HttpServletResponse response, SessionAuthenticat response.addHeader(HttpHeaders.SET_COOKIE, buildRefreshCookie(sessionAuthentication.refreshToken()).toString()); } + public void writeImpersonationCookie(HttpServletResponse response, String cookieName, long targetUserId) { + response.addHeader(HttpHeaders.SET_COOKIE, ResponseCookie.from(cookieName, Long.toString(targetUserId)) + .httpOnly(true) + .secure(cookieSecure) + .sameSite(cookieSameSite) + .path("/") + .maxAge(Duration.ofDays(refreshDays)) + .build() + .toString()); + } + + public void clearImpersonationCookie(HttpServletResponse response, String cookieName) { + response.addHeader(HttpHeaders.SET_COOKIE, clearCookie(cookieName).toString()); + } + + public void clearImpersonationCookie(HttpServletResponse response) { + clearImpersonationCookie(response, impersonateCookieName); + } + public void clearSessionCookies(HttpServletResponse response) { response.addHeader(HttpHeaders.SET_COOKIE, clearCookie(accessCookieName).toString()); response.addHeader(HttpHeaders.SET_COOKIE, clearCookie(refreshCookieName).toString()); + response.addHeader(HttpHeaders.SET_COOKIE, clearCookie(impersonateCookieName).toString()); } private SessionAuthentication rotateSession( diff --git a/backend/src/main/java/com/dbaagent/service/ConnectionChatAccessPolicyService.java b/backend/src/main/java/com/dbaagent/service/ConnectionChatAccessPolicyService.java index d657e80..d3175e4 100644 --- a/backend/src/main/java/com/dbaagent/service/ConnectionChatAccessPolicyService.java +++ b/backend/src/main/java/com/dbaagent/service/ConnectionChatAccessPolicyService.java @@ -3,6 +3,7 @@ import com.dbaagent.dto.ConnectionChatAccessPolicyResponse; import com.dbaagent.dto.PolicyPreviewResponse; import com.dbaagent.model.ConnectionChatAccessPolicy; +import com.dbaagent.model.ColumnMetadata; import com.dbaagent.model.SecurityEventOutcome; import com.dbaagent.model.SecurityEventType; import com.dbaagent.model.SchemaMetadata; @@ -24,11 +25,46 @@ import java.util.Map; import java.util.Optional; import java.util.Set; +import java.util.regex.Matcher; +import java.util.regex.Pattern; @Service @RequiredArgsConstructor public class ConnectionChatAccessPolicyService { + private static final Pattern ONLY_SCHEMA_PATTERN = Pattern.compile( + "(?:only|just)\\s+(?:have\\s+)?(?:access\\s+to\\s+)?(?:the\\s+)?schema\\s+([a-z_][a-z0-9_]*)", + Pattern.CASE_INSENSITIVE + ); + private static final Pattern ACCESS_ONLY_SCHEMA_PATTERN = Pattern.compile( + "access\\s+only\\s+to\\s+(?:schema\\s+)?([a-z_][a-z0-9_]*)", + Pattern.CASE_INSENSITIVE + ); + // Prefix-only: remainder is sliced linearly in extractConstraints so we never + // run `.+?` + `\s+` lookaheads on untrusted policy text (ReDoS / CodeQL). + private static final Pattern DENY_PREFIX_PATTERN = Pattern.compile( + "(?:cannot|can't|must not|do not|don't|never) (?:query|access|see|select|read|use|return|expose) " + + "|(?:redact|block|deny|hide) ", + Pattern.CASE_INSENSITIVE + ); + private static final Pattern ALLOW_PREFIX_PATTERN = Pattern.compile( + "(?:but |except(?: that)? )?\\bcan (?:query|access|see|select|read|use|return) ", + Pattern.CASE_INSENSITIVE + ); + private static final Pattern DENY_STOP_PATTERN = Pattern.compile( + " (?:but|except|however|strictly)\\b|[.;]" + ); + private static final Pattern ALLOW_STOP_PATTERN = Pattern.compile( + " (?:strictly)\\b|[.;]" + ); + private static final List TYPE_FAMILIES = List.of( + new TypeFamily("integer", Set.of("int", "integer", "bigint", "smallint", "tinyint", "serial", "bigserial", "int2", "int4", "int8")), + new TypeFamily("float", Set.of("float", "double", "real", "numeric", "decimal", "number", "money", "float4", "float8")), + new TypeFamily("string", Set.of("varchar", "character varying", "character", "char", "text", "clob", "uuid", "json", "jsonb", "string", "citext")), + new TypeFamily("boolean", Set.of("bool", "boolean")), + new TypeFamily("temporal", Set.of("date", "time", "timestamp", "timestamptz", "datetime", "interval")) + ); + private final ConnectionChatAccessPolicyRepository policyRepository; private final TableClassificationRepository tableClassificationRepository; private final ObjectProvider schemaScannerServiceProvider; @@ -89,7 +125,8 @@ public ConnectionChatAccessPolicyResponse savePolicy( "updatedBy", updatedBy, "blockedSensitivityCategories", parsedPolicy.blockedSensitivityCategories(), "deniedTables", parsedPolicy.deniedTables(), - "deniedColumns", parsedPolicy.deniedColumns() + "deniedColumns", parsedPolicy.deniedColumns(), + "allowedSchemas", parsedPolicy.allowedSchemas() )) .build()); @@ -140,6 +177,7 @@ private EffectivePolicy toEffectivePolicy(ConnectionChatAccessPolicy policy) { new LinkedHashSet<>(policy.getBlockedSensitivityCategories() == null ? List.of() : policy.getBlockedSensitivityCategories()), new LinkedHashSet<>(policy.getDeniedTables() == null ? List.of() : policy.getDeniedTables()), new LinkedHashSet<>(policy.getDeniedColumns() == null ? List.of() : policy.getDeniedColumns()), + new LinkedHashSet<>(parsedPolicy.allowedSchemas()), policy.isBlockMode(), policy.isRedactMode(), policy.getPlainEnglishPolicy(), @@ -149,29 +187,7 @@ private EffectivePolicy toEffectivePolicy(ConnectionChatAccessPolicy policy) { } private ParsedPolicy parsedFromPolicy(ConnectionChatAccessPolicy policy) { - List impactedTables = new ArrayList<>(); - List impactedColumns = new ArrayList<>(); - Map descriptors = buildProtectionDescriptors( - policy.getConnectionId(), - policy.getBlockedSensitivityCategories(), - policy.getDeniedTables(), - policy.getDeniedColumns() - ); - descriptors.values().forEach(descriptor -> { - if (descriptor.protectWholeTable()) { - impactedTables.add(descriptor.tableName()); - } - descriptor.restrictedColumns().forEach(column -> impactedColumns.add(descriptor.tableName() + "." + column)); - }); - return new ParsedPolicy( - normalizeList(policy.getBlockedSensitivityCategories()), - normalizeList(policy.getDeniedTables()), - normalizeList(policy.getDeniedColumns()), - impactedTables, - impactedColumns, - policy.isBlockMode(), - policy.isRedactMode() - ); + return parsePolicy(policy.getConnectionId(), policy.getPlainEnglishPolicy()); } private ParsedPolicy parsePolicy(String connectionId, String plainEnglishPolicy) { @@ -202,39 +218,49 @@ private ParsedPolicy parsePolicy(String connectionId, String plainEnglishPolicy) } SchemaMetadata schemaMetadata = tryScanSchema(connectionId); + Set allowedSchemas = extractAllowedSchemas(normalized, schemaMetadata); LinkedHashSet deniedTables = new LinkedHashSet<>(); LinkedHashSet deniedColumns = new LinkedHashSet<>(); + if (schemaMetadata != null) { for (TableMetadata table : schemaMetadata.getTables()) { - if (containsWord(normalized, table.getName())) { - deniedTables.add(table.getName()); + if (!schemaInScope(table.getSchema(), allowedSchemas)) { + continue; + } + String qualifiedTable = qualifyTable(table.getSchema(), table.getName()); + if (containsWord(normalized, qualifiedTable) || containsWord(normalized, table.getName())) { + deniedTables.add(qualifiedTable); } if (table.getColumns() != null) { table.getColumns().forEach(column -> { - String qualified = table.getName() + "." + column.getName(); - if (containsWord(normalized, qualified) || containsWord(normalized, column.getName())) { - deniedColumns.add(qualified); + String qualifiedColumn = qualifiedTable + "." + column.getName(); + if (containsWord(normalized, qualifiedColumn)) { + deniedColumns.add(qualifiedColumn); } }); } } + applyColumnConstraints(normalized, allowedSchemas, schemaMetadata, deniedColumns); } Map descriptors = buildProtectionDescriptors( connectionId, new ArrayList<>(blockedCategories), new ArrayList<>(deniedTables), - new ArrayList<>(deniedColumns) + new ArrayList<>(deniedColumns), + allowedSchemas, + schemaMetadata ); List impactedTables = descriptors.values().stream() .filter(ProtectionDescriptor::protectWholeTable) - .map(ProtectionDescriptor::tableName) + .map(ProtectionDescriptor::qualifiedTableName) .distinct() .sorted() .toList(); List impactedColumns = descriptors.values().stream() - .flatMap(descriptor -> descriptor.restrictedColumns().stream().map(column -> descriptor.tableName() + "." + column)) + .flatMap(descriptor -> descriptor.restrictedColumns().stream() + .map(column -> descriptor.qualifiedTableName() + "." + column)) .distinct() .sorted() .toList(); @@ -243,6 +269,7 @@ private ParsedPolicy parsePolicy(String connectionId, String plainEnglishPolicy) new ArrayList<>(blockedCategories), new ArrayList<>(deniedTables), new ArrayList<>(deniedColumns), + new ArrayList<>(allowedSchemas), impactedTables, impactedColumns, true, @@ -256,52 +283,353 @@ public Map buildProtectionDescriptors( List deniedTables, List deniedColumns ) { + return buildProtectionDescriptors( + connectionId, + blockedSensitivityCategories, + deniedTables, + deniedColumns, + Set.of(), + tryScanSchema(connectionId) + ); + } + + public Map buildProtectionDescriptors( + String connectionId, + List blockedSensitivityCategories, + List deniedTables, + List deniedColumns, + Set allowedSchemas, + SchemaMetadata schemaMetadata + ) { + SchemaMetadata effectiveSchema = schemaMetadata != null ? schemaMetadata : tryScanSchema(connectionId); Set categorySet = normalizeSet(blockedSensitivityCategories); Set deniedTableSet = normalizeSet(deniedTables); Set deniedColumnSet = normalizeSet(deniedColumns); + Set allowedSchemaSet = normalizeSet(new ArrayList<>(allowedSchemas)); + + Map> schemasByBareTable = indexSchemasByBareTable(effectiveSchema); Map descriptors = new LinkedHashMap<>(); for (TableClassification classification : tableClassificationRepository.findLatestByConnectionIdOrderByTableNameAsc(connectionId)) { - String tableName = classification.getTableName(); - ProtectionDescriptor descriptor = descriptors.computeIfAbsent( - normalizeName(tableName), - ignored -> new ProtectionDescriptor(tableName, false, new LinkedHashSet<>()) - ); - - boolean tableExplicitlyDenied = deniedTableSet.contains(normalizeName(tableName)); - boolean tableCategoryBlocked = categorySet.contains(normalizeName(classification.getSensitivityLevel())); - if (tableExplicitlyDenied) { - descriptor.protectWholeTable = true; - } - if (tableCategoryBlocked && (classification.getSensitiveColumns() == null || classification.getSensitiveColumns().isEmpty())) { - descriptor.protectWholeTable = true; - } + String bareTable = classification.getTableName(); + List schemas = schemasByBareTable.getOrDefault(normalizeName(bareTable), List.of("")); + for (String schema : schemas) { + if (!schemaInScope(schema, allowedSchemaSet)) { + continue; + } + String qualifiedTable = qualifyTable(schema, bareTable); + ProtectionDescriptor descriptor = descriptors.computeIfAbsent( + normalizeName(qualifiedTable), + ignored -> new ProtectionDescriptor(schema, bareTable, false, new LinkedHashSet<>()) + ); + + boolean tableExplicitlyDenied = deniedTableSet.contains(normalizeName(qualifiedTable)) + || deniedTableSet.contains(normalizeName(bareTable)); + boolean tableCategoryBlocked = categorySet.contains(normalizeName(classification.getSensitivityLevel())); + if (tableExplicitlyDenied) { + descriptor.protectWholeTable = true; + } + if (tableCategoryBlocked && (classification.getSensitiveColumns() == null || classification.getSensitiveColumns().isEmpty())) { + descriptor.protectWholeTable = true; + } - if (classification.getSensitiveColumns() != null) { - for (Map sensitiveColumn : classification.getSensitiveColumns()) { - String column = String.valueOf(sensitiveColumn.get("column")); - String type = sensitiveColumn.get("type") == null ? "" : String.valueOf(sensitiveColumn.get("type")); - if (categorySet.contains(normalizeName(type)) - || deniedColumnSet.contains(normalizeName(tableName + "." + column)) - || deniedColumnSet.contains(normalizeName(column)) - || tableExplicitlyDenied) { - descriptor.restrictedColumns.add(column); + if (classification.getSensitiveColumns() != null) { + for (Map sensitiveColumn : classification.getSensitiveColumns()) { + String column = String.valueOf(sensitiveColumn.get("column")); + String type = sensitiveColumn.get("type") == null ? "" : String.valueOf(sensitiveColumn.get("type")); + String qualifiedColumn = qualifiedTable + "." + column; + if (categorySet.contains(normalizeName(type)) + || deniedColumnSet.contains(normalizeName(qualifiedColumn)) + || tableExplicitlyDenied) { + descriptor.restrictedColumns.add(column); + } } } } } - deniedTableSet.forEach(table -> descriptors.computeIfAbsent(table, key -> new ProtectionDescriptor(table, true, new LinkedHashSet<>())).protectWholeTable = true); + deniedTableSet.forEach(tableRef -> { + String normalized = normalizeName(tableRef); + descriptors.computeIfAbsent(normalized, key -> descriptorFromTableRef(tableRef, true)).protectWholeTable = true; + }); deniedColumnSet.forEach(columnRef -> { - String[] parts = columnRef.split("\\.", 2); - if (parts.length == 2) { - ProtectionDescriptor descriptor = descriptors.computeIfAbsent(parts[0], key -> new ProtectionDescriptor(parts[0], false, new LinkedHashSet<>())); + String[] parts = columnRef.split("\\.", 3); + if (parts.length == 3) { + String qualifiedTable = parts[0] + "." + parts[1]; + ProtectionDescriptor descriptor = descriptors.computeIfAbsent( + normalizeName(qualifiedTable), + key -> descriptorFromTableRef(qualifiedTable, false) + ); + descriptor.restrictedColumns.add(parts[2]); + } else if (parts.length == 2) { + ProtectionDescriptor descriptor = descriptors.computeIfAbsent( + normalizeName(columnRef.substring(0, columnRef.lastIndexOf('.'))), + key -> descriptorFromTableRef(parts[0], false) + ); descriptor.restrictedColumns.add(parts[1]); } }); return descriptors; } + public Map buildProtectionDescriptors( + ConnectionChatAccessPolicyService.EffectivePolicy policy + ) { + if (policy == null || !policy.protectsAnything()) { + return Map.of(); + } + return buildProtectionDescriptors( + policy.connectionId(), + new ArrayList<>(policy.blockedSensitivityCategories()), + new ArrayList<>(policy.deniedTables()), + new ArrayList<>(policy.deniedColumns()), + policy.allowedSchemas(), + null + ); + } + + static String qualifyTable(String schema, String table) { + if (table == null || table.isBlank()) { + return ""; + } + if (schema == null || schema.isBlank() || "public".equalsIgnoreCase(schema)) { + return table; + } + return schema + "." + table; + } + + private ProtectionDescriptor descriptorFromTableRef(String tableRef, boolean protectWholeTable) { + String[] parts = tableRef.split("\\.", 2); + if (parts.length == 2) { + return new ProtectionDescriptor(parts[0], parts[1], protectWholeTable, new LinkedHashSet<>()); + } + return new ProtectionDescriptor(null, tableRef, protectWholeTable, new LinkedHashSet<>()); + } + + private Map> indexSchemasByBareTable(SchemaMetadata schemaMetadata) { + Map> schemasByBareTable = new LinkedHashMap<>(); + if (schemaMetadata == null || schemaMetadata.getTables() == null) { + return schemasByBareTable; + } + for (TableMetadata table : schemaMetadata.getTables()) { + schemasByBareTable + .computeIfAbsent(normalizeName(table.getName()), ignored -> new ArrayList<>()) + .add(table.getSchema() == null ? "" : table.getSchema()); + } + return schemasByBareTable; + } + + private void applyColumnConstraints( + String normalized, + Set allowedSchemas, + SchemaMetadata schemaMetadata, + Set deniedColumns + ) { + Set knownColumnNames = collectColumnNames(schemaMetadata, allowedSchemas); + List denials = extractConstraints(normalized, DENY_PREFIX_PATTERN, DENY_STOP_PATTERN, knownColumnNames); + List allowances = extractConstraints(normalized, ALLOW_PREFIX_PATTERN, ALLOW_STOP_PATTERN, knownColumnNames); + if (denials.isEmpty()) { + return; + } + + for (TableMetadata table : schemaMetadata.getTables()) { + if (!schemaInScope(table.getSchema(), allowedSchemas) || table.getColumns() == null) { + continue; + } + String qualifiedTable = qualifyTable(table.getSchema(), table.getName()); + for (ColumnMetadata column : table.getColumns()) { + boolean denied = denials.stream().anyMatch(constraint -> constraint.matches(column)); + boolean allowed = allowances.stream().anyMatch(constraint -> constraint.matches(column)); + if (denied && !allowed) { + deniedColumns.add(qualifiedTable + "." + column.getName()); + } + } + } + } + + private List extractConstraints( + String normalized, + Pattern prefixPattern, + Pattern stopPattern, + Set knownColumnNames + ) { + List constraints = new ArrayList<>(); + String haystack = normalized == null ? "" : normalized.replaceAll("\\s+", " ").trim(); + Matcher matcher = prefixPattern.matcher(haystack); + while (matcher.find()) { + String snippet = sliceUntilStop(haystack, matcher.end(), stopPattern); + if (snippet == null || snippet.isBlank()) { + continue; + } + String lowered = snippet.toLowerCase(Locale.ROOT); + LinkedHashSet typeKeys = new LinkedHashSet<>(); + for (TypeFamily family : TYPE_FAMILIES) { + if (family.mentionedIn(lowered)) { + typeKeys.add(family.key()); + } + } + LinkedHashSet nameTokens = new LinkedHashSet<>(); + knownColumnNames.stream() + .sorted((left, right) -> Integer.compare(right.length(), left.length())) + .filter(name -> containsWholeWord(lowered, name)) + .forEach(nameTokens::add); + if (!typeKeys.isEmpty() || !nameTokens.isEmpty()) { + constraints.add(new ColumnConstraint(typeKeys, nameTokens)); + } + } + return constraints; + } + + private Set collectColumnNames(SchemaMetadata schemaMetadata, Set allowedSchemas) { + LinkedHashSet names = new LinkedHashSet<>(); + if (schemaMetadata == null || schemaMetadata.getTables() == null) { + return names; + } + for (TableMetadata table : schemaMetadata.getTables()) { + if (!schemaInScope(table.getSchema(), allowedSchemas) || table.getColumns() == null) { + continue; + } + for (ColumnMetadata column : table.getColumns()) { + String name = normalizeName(column.getName()); + if (!name.isBlank() && !isTypeToken(name)) { + names.add(name); + } + } + } + return names; + } + + private boolean isTypeToken(String name) { + return TYPE_FAMILIES.stream().anyMatch(family -> family.aliases().contains(name) || family.key().equals(name)); + } + + private String sliceUntilStop(String haystack, int start, Pattern stopPattern) { + if (start >= haystack.length()) { + return ""; + } + Matcher stop = stopPattern.matcher(haystack); + if (stop.find(start)) { + return haystack.substring(start, stop.start()).trim(); + } + return haystack.substring(start).trim(); + } + + private boolean containsWholeWord(String haystack, String needle) { + if (haystack == null || needle == null || needle.isBlank()) { + return false; + } + return Pattern.compile("\\b" + Pattern.quote(needle) + "\\b", Pattern.CASE_INSENSITIVE) + .matcher(haystack) + .find(); + } + + private record TypeFamily(String key, Set aliases) { + boolean mentionedIn(String snippet) { + if (containsWholeWordStatic(snippet, key)) { + return true; + } + return aliases.stream().anyMatch(alias -> containsWholeWordStatic(snippet, alias)); + } + + boolean matchesDataType(String dataType) { + String normalized = dataType == null + ? "" + : dataType.toLowerCase(Locale.ROOT).replaceAll("\\([^)]*\\)", " ").trim(); + if (normalized.isBlank()) { + return false; + } + if (containsWholeWordStatic(normalized, key) || normalized.equals(key)) { + return true; + } + return aliases.stream().anyMatch(alias -> + containsWholeWordStatic(normalized, alias) || normalized.equals(alias) + ); + } + + private static boolean containsWholeWordStatic(String haystack, String needle) { + if (haystack == null || needle == null || needle.isBlank()) { + return false; + } + return Pattern.compile("\\b" + Pattern.quote(needle) + "\\b", Pattern.CASE_INSENSITIVE) + .matcher(haystack) + .find(); + } + } + + private record ColumnConstraint(Set typeKeys, Set nameTokens) { + boolean matches(ColumnMetadata column) { + if ((typeKeys == null || typeKeys.isEmpty()) && (nameTokens == null || nameTokens.isEmpty())) { + return false; + } + boolean typeOk = typeKeys == null || typeKeys.isEmpty() + || typeKeys.stream().anyMatch(key -> TYPE_FAMILIES.stream() + .filter(family -> family.key().equals(key)) + .anyMatch(family -> family.matchesDataType(column.getDataType()))); + boolean nameOk = nameTokens == null || nameTokens.isEmpty() + || nameTokens.stream().anyMatch(token -> columnNameMatches(column.getName(), token)); + return typeOk && nameOk; + } + + private static boolean columnNameMatches(String columnName, String token) { + String column = columnName == null ? "" : columnName.trim().replace("\"", "").replace("`", "").toLowerCase(Locale.ROOT); + String needle = token == null ? "" : token.trim().toLowerCase(Locale.ROOT); + if (column.isBlank() || needle.isBlank()) { + return false; + } + if (column.equals(needle)) { + return true; + } + return column.startsWith(needle + "_") + || column.endsWith("_" + needle) + || column.contains("_" + needle + "_"); + } + } + + private Set extractAllowedSchemas(String normalized, SchemaMetadata schemaMetadata) { + LinkedHashSet schemas = new LinkedHashSet<>(); + for (Pattern pattern : List.of(ONLY_SCHEMA_PATTERN, ACCESS_ONLY_SCHEMA_PATTERN)) { + Matcher matcher = pattern.matcher(normalized); + while (matcher.find()) { + schemas.add(normalizeName(matcher.group(1))); + } + } + if (schemas.isEmpty() && containsAny(normalized, "cannot access any other schema", "no other schema")) { + for (Pattern pattern : List.of( + Pattern.compile("schema\\s+([a-z_][a-z0-9_]*)", Pattern.CASE_INSENSITIVE) + )) { + Matcher matcher = pattern.matcher(normalized); + while (matcher.find()) { + schemas.add(normalizeName(matcher.group(1))); + } + } + } + if (schemaMetadata != null && !schemas.isEmpty()) { + Set known = new LinkedHashSet<>(); + for (TableMetadata table : schemaMetadata.getTables()) { + if (table.getSchema() != null) { + known.add(normalizeName(table.getSchema())); + } + } + schemas.retainAll(known); + } + return schemas; + } + + private boolean schemaInScope(String schema, Set allowedSchemas) { + return isSchemaInScope(schema, allowedSchemas); + } + + public static boolean isSchemaInScope(String schema, Set allowedSchemas) { + if (allowedSchemas == null || allowedSchemas.isEmpty()) { + return true; + } + String normalized = schema == null ? "" : schema.trim().replace("\"", "").replace("`", "").toLowerCase(Locale.ROOT); + if (normalized.isBlank()) { + normalized = "public"; + } + return allowedSchemas.contains(normalized); + } + private SchemaMetadata tryScanSchema(String connectionId) { try { SchemaScannerService schemaScannerService = schemaScannerServiceProvider.getIfAvailable(); @@ -348,7 +676,7 @@ private List normalizeList(List values) { } private String normalizeName(String value) { - return value == null ? "" : value.trim().toLowerCase(Locale.ROOT); + return value == null ? "" : value.trim().replace("\"", "").replace("`", "").toLowerCase(Locale.ROOT); } public record EffectivePolicy( @@ -358,6 +686,7 @@ public record EffectivePolicy( Set blockedSensitivityCategories, Set deniedTables, Set deniedColumns, + Set allowedSchemas, boolean blockMode, boolean redactMode, String plainEnglishPolicy, @@ -365,11 +694,16 @@ public record EffectivePolicy( List impactedColumns ) { public static EffectivePolicy none() { - return new EffectivePolicy(false, null, null, Set.of(), Set.of(), Set.of(), false, false, null, List.of(), List.of()); + return new EffectivePolicy(false, null, null, Set.of(), Set.of(), Set.of(), Set.of(), false, false, null, List.of(), List.of()); } public boolean protectsAnything() { - return present && (!blockedSensitivityCategories.isEmpty() || !deniedTables.isEmpty() || !deniedColumns.isEmpty()); + return present && ( + !blockedSensitivityCategories.isEmpty() + || !deniedTables.isEmpty() + || !deniedColumns.isEmpty() + || !allowedSchemas.isEmpty() + ); } } @@ -377,6 +711,7 @@ private record ParsedPolicy( List blockedSensitivityCategories, List deniedTables, List deniedColumns, + List allowedSchemas, List impactedTables, List impactedColumns, boolean blockMode, @@ -385,20 +720,30 @@ private record ParsedPolicy( } public static final class ProtectionDescriptor { + private final String schemaName; private final String tableName; private boolean protectWholeTable; private final Set restrictedColumns; - private ProtectionDescriptor(String tableName, boolean protectWholeTable, Set restrictedColumns) { + private ProtectionDescriptor(String schemaName, String tableName, boolean protectWholeTable, Set restrictedColumns) { + this.schemaName = schemaName; this.tableName = tableName; this.protectWholeTable = protectWholeTable; this.restrictedColumns = restrictedColumns; } + public String schemaName() { + return schemaName; + } + public String tableName() { return tableName; } + public String qualifiedTableName() { + return qualifyTable(schemaName, tableName); + } + public boolean protectWholeTable() { return protectWholeTable; } diff --git a/backend/src/main/java/com/dbaagent/service/ImpersonationService.java b/backend/src/main/java/com/dbaagent/service/ImpersonationService.java new file mode 100644 index 0000000..2d9750a --- /dev/null +++ b/backend/src/main/java/com/dbaagent/service/ImpersonationService.java @@ -0,0 +1,327 @@ +package com.dbaagent.service; + +import com.dbaagent.model.SecurityEventOutcome; +import com.dbaagent.model.SecurityEventType; +import com.dbaagent.model.User; +import com.dbaagent.model.UserAccountStatus; +import com.dbaagent.repository.UserRepository; +import com.dbaagent.security.CustomUserDetailsService; +import com.dbaagent.security.ImpersonationContext; +import jakarta.servlet.http.Cookie; +import jakarta.servlet.http.HttpServletRequest; +import jakarta.servlet.http.HttpServletResponse; +import lombok.RequiredArgsConstructor; +import lombok.extern.slf4j.Slf4j; +import org.springframework.beans.factory.annotation.Value; +import org.springframework.http.HttpHeaders; +import org.springframework.security.authentication.UsernamePasswordAuthenticationToken; +import org.springframework.security.core.Authentication; +import org.springframework.security.core.GrantedAuthority; +import org.springframework.security.core.context.SecurityContextHolder; +import org.springframework.security.core.userdetails.UserDetails; +import org.springframework.security.core.userdetails.UsernameNotFoundException; +import org.springframework.stereotype.Service; +import org.springframework.web.server.ResponseStatusException; + +import java.util.LinkedHashMap; +import java.util.List; +import java.util.Map; +import java.util.Optional; + +import static org.springframework.http.HttpStatus.BAD_REQUEST; +import static org.springframework.http.HttpStatus.FORBIDDEN; +import static org.springframework.http.HttpStatus.NOT_FOUND; + +/** + * Admin-only profile switch. The admin session (JWT cookies) is unchanged; + * a separate httpOnly cookie names the user to evaluate as. The JWT filter + * overlays that principal onto the SecurityContext for every request except + * the impersonation control plane, logout, and session refresh. + */ +@Service +@RequiredArgsConstructor +@Slf4j +public class ImpersonationService { + + static final String DEFAULT_COOKIE_NAME = "impersonate_user"; + + private final UserRepository userRepository; + private final CustomUserDetailsService userDetailsService; + private final AuthSessionService authSessionService; + private final SecurityEventService securityEventService; + + @Value("${security.auth.enabled:true}") + private boolean authEnabled; + + @Value("${security.cookie.impersonate-name:" + DEFAULT_COOKIE_NAME + "}") + private String impersonateCookieName; + + public ImpersonationContext.State start( + User actor, + Long targetUserId, + HttpServletRequest request, + HttpServletResponse response + ) { + requireAdminActor(actor); + User target = requireAllowedTarget(actor, targetUserId); + authSessionService.writeImpersonationCookie(response, impersonateCookieName, target.getId()); + securityEventService.log(SecurityEventService.EventRequest.builder() + .eventType(SecurityEventType.IMPERSONATION_STARTED) + .outcome(SecurityEventOutcome.SUCCESS) + .userId(target.getId()) + .actorUserId(actor.getId()) + .email(actor.getEmail()) + .targetResource("user:" + target.getId()) + .clientIp(clientIp(request)) + .userAgent(userAgent(request)) + .metadata(Map.of( + "impersonatorUsername", actor.getUsername(), + "targetUsername", target.getUsername() + )) + .build()); + log.info("Admin {} started profile switch to {}", actor.getUsername(), target.getUsername()); + return new ImpersonationContext.State(actor, target); + } + + public User stop( + User actor, + HttpServletRequest request, + HttpServletResponse response + ) { + requireAdminActor(actor); + Optional target = readTargetUser(request); + authSessionService.clearImpersonationCookie(response, impersonateCookieName); + ImpersonationContext.clear(); + target.ifPresent(stopped -> securityEventService.log(SecurityEventService.EventRequest.builder() + .eventType(SecurityEventType.IMPERSONATION_STOPPED) + .outcome(SecurityEventOutcome.SUCCESS) + .userId(stopped.getId()) + .actorUserId(actor.getId()) + .email(actor.getEmail()) + .targetResource("user:" + stopped.getId()) + .clientIp(clientIp(request)) + .userAgent(userAgent(request)) + .metadata(Map.of( + "impersonatorUsername", actor.getUsername(), + "targetUsername", stopped.getUsername() + )) + .build())); + log.info("Admin {} stopped profile switch", actor.getUsername()); + return actor; + } + + public List> listCandidates(User actor) { + requireAdminActor(actor); + return userRepository.findAll().stream() + .filter(user -> isAllowedTarget(actor, user)) + .map(this::toCandidate) + .toList(); + } + + public Optional resolveFromCookie(HttpServletRequest request, User sessionUser) { + if (sessionUser == null || !sessionUser.isAdmin()) { + return Optional.empty(); + } + return readTargetUser(request) + .filter(target -> isAllowedTarget(sessionUser, target)) + .map(target -> new ImpersonationContext.State(sessionUser, target)); + } + + public void decorateAuthPayload(HttpServletRequest request, User sessionUser, Map payload) { + Optional state = ImpersonationContext.current(); + if (state.isEmpty()) { + state = resolveFromCookie(request, sessionUser); + } + if (state.isEmpty()) { + payload.put("impersonating", false); + return; + } + payload.put("impersonating", true); + payload.put("impersonatorUsername", state.get().impersonatorUsername()); + payload.put("impersonatorEmail", state.get().impersonatorEmail()); + } + + /** + * Overlay the target principal when the admin JWT (or the auth-disabled + * synthetic admin) is already in the SecurityContext. No-ops on the + * impersonation control-plane, logout/refresh, MCP tokens, and invalid cookies. + */ + public void applyToRequest(HttpServletRequest request) { + if (!shouldApply(request)) { + return; + } + Authentication current = SecurityContextHolder.getContext().getAuthentication(); + if (current == null || !current.isAuthenticated() || "anonymousUser".equals(current.getPrincipal())) { + return; + } + if (authEnabled && !hasAdminRole(current)) { + return; + } + User impersonator = userRepository.findByUsername(current.getName()) + .orElseGet(() -> syntheticAdmin(current.getName())); + if (!impersonator.isAdmin() && authEnabled) { + return; + } + Optional target = readTargetUser(request); + if (target.isEmpty() || !isAllowedTarget(impersonator, target.get())) { + return; + } + UserDetails details; + try { + details = userDetailsService.loadUserByUsername(target.get().getUsername()); + } catch (UsernameNotFoundException e) { + return; + } + UsernamePasswordAuthenticationToken swapped = new UsernamePasswordAuthenticationToken( + details, + null, + details.getAuthorities() + ); + swapped.setDetails(current.getDetails()); + SecurityContextHolder.getContext().setAuthentication(swapped); + ImpersonationContext.enter(new ImpersonationContext.State(impersonator, target.get())); + log.debug("Applied profile switch: {} -> {}", impersonator.getUsername(), target.get().getUsername()); + } + + boolean shouldApply(HttpServletRequest request) { + String path = request.getServletPath() != null ? request.getServletPath() : ""; + String uri = request.getRequestURI() != null ? request.getRequestURI() : ""; + if (isControlPlane(path) || isControlPlane(uri)) { + return false; + } + String authorization = request.getHeader(HttpHeaders.AUTHORIZATION); + if (authorization != null && authorization.startsWith("Bearer ")) { + String token = authorization.substring(7); + if (token.startsWith(McpTokenService.TOKEN_PREFIX)) { + return false; + } + } + return true; + } + + private boolean isControlPlane(String path) { + if (path == null || path.isBlank()) { + return false; + } + return path.contains("/admin/impersonate") + || path.endsWith("/auth/logout") + || path.endsWith("/auth/logout-all") + || path.endsWith("/auth/refresh"); + } + + private Optional readTargetUser(HttpServletRequest request) { + Long userId = readTargetUserId(request); + if (userId == null) { + return Optional.empty(); + } + return userRepository.findById(userId); + } + + private Long readTargetUserId(HttpServletRequest request) { + Cookie[] cookies = request.getCookies(); + if (cookies == null) { + return null; + } + for (Cookie cookie : cookies) { + if (impersonateCookieName.equals(cookie.getName())) { + return parseUserId(cookie.getValue()); + } + } + return null; + } + + private Long parseUserId(String value) { + if (value == null || value.isBlank()) { + return null; + } + try { + long parsed = Long.parseLong(value.trim()); + return parsed > 0 ? parsed : null; + } catch (NumberFormatException e) { + return null; + } + } + + private void requireAdminActor(User actor) { + if (actor == null || !actor.isAdmin()) { + throw new ResponseStatusException(FORBIDDEN, "Only administrators can switch profiles"); + } + } + + private User requireAllowedTarget(User actor, Long targetUserId) { + if (targetUserId == null) { + throw new ResponseStatusException(BAD_REQUEST, "userId is required"); + } + User target = userRepository.findById(targetUserId) + .orElseThrow(() -> new ResponseStatusException(NOT_FOUND, "User not found")); + if (!isAllowedTarget(actor, target)) { + throw new ResponseStatusException(BAD_REQUEST, denialReason(actor, target)); + } + return target; + } + + boolean isAllowedTarget(User actor, User target) { + if (actor == null || target == null || target.getId() == null) { + return false; + } + if (actor.getId() != null && actor.getId().equals(target.getId())) { + return false; + } + if (target.isAdmin()) { + return false; + } + return target.getAccountStatusEnum() == UserAccountStatus.ACTIVE; + } + + private String denialReason(User actor, User target) { + if (actor.getId() != null && actor.getId().equals(target.getId())) { + return "Cannot switch into your own profile"; + } + if (target.isAdmin()) { + return "Cannot switch into another administrator profile"; + } + if (target.getAccountStatusEnum() != UserAccountStatus.ACTIVE) { + return "Cannot switch into a locked or disabled account"; + } + return "Cannot switch into this profile"; + } + + private Map toCandidate(User user) { + Map dto = new LinkedHashMap<>(); + dto.put("id", user.getId()); + dto.put("username", user.getUsername()); + dto.put("email", user.getEmail()); + dto.put("role", user.getRole()); + dto.put("accountStatus", user.getAccountStatus()); + return dto; + } + + private boolean hasAdminRole(Authentication authentication) { + return authentication.getAuthorities().stream() + .map(GrantedAuthority::getAuthority) + .anyMatch("ROLE_ADMIN"::equals); + } + + private User syntheticAdmin(String username) { + User admin = new User(); + admin.setUsername(username != null && !username.isBlank() ? username : "admin"); + admin.setRole("ADMIN"); + admin.setAccountStatus(UserAccountStatus.ACTIVE.name()); + return admin; + } + + private String clientIp(HttpServletRequest request) { + if (request == null) { + return null; + } + String forwarded = request.getHeader("X-Forwarded-For"); + if (forwarded != null && !forwarded.isBlank()) { + return forwarded.split(",")[0].trim(); + } + return request.getRemoteAddr(); + } + + private String userAgent(HttpServletRequest request) { + return request == null ? null : request.getHeader("User-Agent"); + } +} diff --git a/backend/src/main/java/com/dbaagent/service/UserDataAccessPolicyService.java b/backend/src/main/java/com/dbaagent/service/UserDataAccessPolicyService.java index a65a9f6..82f7632 100644 --- a/backend/src/main/java/com/dbaagent/service/UserDataAccessPolicyService.java +++ b/backend/src/main/java/com/dbaagent/service/UserDataAccessPolicyService.java @@ -2,6 +2,9 @@ import com.dbaagent.model.QueryRequest; import com.dbaagent.model.QueryResult; +import com.dbaagent.model.DatabaseObject; +import com.dbaagent.model.SchemaMetadata; +import com.dbaagent.model.TableMetadata; import com.dbaagent.model.SecurityEventOutcome; import com.dbaagent.model.SecurityEventType; import lombok.RequiredArgsConstructor; @@ -109,16 +112,13 @@ public QueryGuardDecision enforcePreExecution( return QueryGuardDecision.allow(policy); } - Map protectedObjects = policyService.buildProtectionDescriptors( - connectionId, - new ArrayList<>(policy.blockedSensitivityCategories()), - new ArrayList<>(policy.deniedTables()), - new ArrayList<>(policy.deniedColumns()) - ); + Map protectedObjects = + policyService.buildProtectionDescriptors(policy); try { Statement parsed = CCJSqlParserUtil.parse(queryRequest.getQuery()); if (parsed instanceof Select select && select.getPlainSelect() != null) { + enforceAllowedSchemas(select.getPlainSelect(), policy.allowedSchemas()); QueryInspection inspection = inspectPlainSelect(select.getPlainSelect(), protectedObjects); if (inspection.selectsWildcardFromProtectedTable || inspection.rawProtectedColumnsSelected) { logPolicyEvent(SecurityEventType.CHAT_ACCESS_POLICY_BLOCKED, executionContext.actorUsername(), connectionId, "sql_blocked", Map.of( @@ -151,6 +151,96 @@ public QueryGuardDecision enforcePreExecution( return QueryGuardDecision.allow(policy); } + public List filterDatabaseObjects( + String connectionId, + String username, + boolean actorIsAdmin, + List objects + ) { + if (objects == null || objects.isEmpty()) { + return objects; + } + Set allowedSchemas = allowedSchemasForActor(connectionId, username, actorIsAdmin); + if (allowedSchemas.isEmpty()) { + return objects; + } + return objects.stream() + .filter(object -> ConnectionChatAccessPolicyService.isSchemaInScope(object.getSchema(), allowedSchemas)) + .toList(); + } + + public SchemaMetadata filterSchemaMetadata( + String connectionId, + String username, + boolean actorIsAdmin, + SchemaMetadata schema + ) { + if (schema == null) { + return null; + } + Set allowedSchemas = allowedSchemasForActor(connectionId, username, actorIsAdmin); + if (allowedSchemas.isEmpty()) { + return schema; + } + + SchemaMetadata filtered = new SchemaMetadata(); + filtered.setDatabaseName(schema.getDatabaseName()); + filtered.setDbType(schema.getDbType()); + filtered.setTotalViews(schema.getTotalViews()); + filtered.setTotalSizeBytes(schema.getTotalSizeBytes()); + List tables = schema.getTables() == null ? List.of() : schema.getTables().stream() + .filter(table -> ConnectionChatAccessPolicyService.isSchemaInScope(table.getSchema(), allowedSchemas)) + .toList(); + filtered.setTables(tables); + filtered.setTotalTables((long) tables.size()); + if (schema.getRelationships() != null) { + filtered.setRelationships(schema.getRelationships().stream() + .filter(relationship -> + ConnectionChatAccessPolicyService.isSchemaInScope(schemaFromTableRef(relationship.getFromTable()), allowedSchemas) + && ConnectionChatAccessPolicyService.isSchemaInScope(schemaFromTableRef(relationship.getToTable()), allowedSchemas)) + .toList()); + } + return filtered; + } + + public void assertTableSchemaAllowed( + String connectionId, + String username, + boolean actorIsAdmin, + String tableRef + ) { + Set allowedSchemas = allowedSchemasForActor(connectionId, username, actorIsAdmin); + if (allowedSchemas.isEmpty()) { + return; + } + String schema = schemaFromTableRef(tableRef); + if (!ConnectionChatAccessPolicyService.isSchemaInScope(schema, allowedSchemas)) { + throw new UserDataAccessPolicyException( + "This object is in schema '" + (schema.isBlank() ? "public" : schema) + + "' which is outside your allowed schema scope.", + "POLICY_SCHEMA_BLOCKED" + ); + } + } + + private Set allowedSchemasForActor(String connectionId, String username, boolean actorIsAdmin) { + ConnectionChatAccessPolicyService.EffectivePolicy policy = + policyService.resolveEffectivePolicy(connectionId, username, actorIsAdmin); + if (policy == null || policy.allowedSchemas() == null) { + return Set.of(); + } + return policy.allowedSchemas(); + } + + private String schemaFromTableRef(String tableRef) { + if (tableRef == null || tableRef.isBlank()) { + return ""; + } + String normalized = normalizeName(tableRef); + int separator = normalized.lastIndexOf('.'); + return separator > 0 ? normalized.substring(0, separator) : ""; + } + public QueryResult redactResult( String connectionId, QueryResult result, @@ -169,12 +259,8 @@ public QueryResult redactResult( return result; } - Map descriptors = policyService.buildProtectionDescriptors( - connectionId, - new ArrayList<>(policy.blockedSensitivityCategories()), - new ArrayList<>(policy.deniedTables()), - new ArrayList<>(policy.deniedColumns()) - ); + Map descriptors = + policyService.buildProtectionDescriptors(policy); if (result.getColumns() == null || result.getRows() == null) { return result; @@ -182,9 +268,12 @@ public QueryResult redactResult( Set protectedColumnNames = new LinkedHashSet<>(); descriptors.values().forEach(descriptor -> { - protectedColumnNames.addAll(descriptor.restrictedColumns().stream().map(value -> value.toLowerCase(Locale.ROOT)).toList()); + String qualifiedTable = descriptor.qualifiedTableName(); + descriptor.restrictedColumns().forEach(column -> + protectedColumnNames.add(normalizeName(qualifiedTable + "." + column)) + ); if (descriptor.protectWholeTable()) { - protectedColumnNames.add(descriptor.tableName().toLowerCase(Locale.ROOT)); + protectedColumnNames.add(normalizeName(qualifiedTable)); } }); @@ -235,16 +324,53 @@ private boolean mentionsProtectedTables(String normalized, ConnectionChatAccessP return policy.impactedTables().stream().anyMatch(table -> normalized.contains(table.toLowerCase(Locale.ROOT))); } + private void enforceAllowedSchemas(PlainSelect select, Set allowedSchemas) { + if (allowedSchemas == null || allowedSchemas.isEmpty()) { + return; + } + Map aliasToTable = buildAliasMap(select); + Set referencedSchemas = new LinkedHashSet<>(); + collectReferencedSchemas(select.getFromItem(), aliasToTable, referencedSchemas); + if (select.getJoins() != null) { + for (Join join : select.getJoins()) { + collectReferencedSchemas(join.getRightItem(), aliasToTable, referencedSchemas); + } + } + for (String schema : referencedSchemas) { + if (!allowedSchemas.contains(normalizeName(schema))) { + throw new UserDataAccessPolicyException( + "This query references schema '" + schema + "' which is outside your allowed schema scope.", + "POLICY_SCHEMA_BLOCKED" + ); + } + } + } + + private void collectReferencedSchemas(FromItem fromItem, Map aliasToTable, Set schemas) { + if (fromItem instanceof Table table) { + String schema = table.getSchemaName(); + if (schema != null && !schema.isBlank()) { + schemas.add(schema); + return; + } + String qualified = aliasToTable.get(normalizeName(table.getName())); + if (qualified != null && qualified.contains(".")) { + schemas.add(qualified.substring(0, qualified.indexOf('.'))); + } + } + } + private boolean containsDangerousProtectedReference( String normalizedSql, Map protectedObjects ) { for (ConnectionChatAccessPolicyService.ProtectionDescriptor descriptor : protectedObjects.values()) { - if (descriptor.protectWholeTable() && normalizedSql.contains(descriptor.tableName().toLowerCase(Locale.ROOT))) { + String qualified = normalizeName(descriptor.qualifiedTableName()); + if (descriptor.protectWholeTable() && normalizedSql.contains(qualified)) { return true; } for (String column : descriptor.restrictedColumns()) { - if (normalizedSql.contains(column.toLowerCase(Locale.ROOT))) { + if (normalizedSql.contains(qualified + "." + normalizeName(column))) { return true; } } @@ -266,18 +392,22 @@ private QueryInspection inspectPlainSelect( select.getSelectItems().forEach(item -> { Expression expression = item.getExpression(); if (expression instanceof AllColumns) { - inspection.selectsWildcardFromProtectedTable = !protectedObjects.isEmpty(); - inspection.reason = "SELECT * touches protected objects"; - inspection.protectedTables.addAll(protectedObjects.values().stream().map(ConnectionChatAccessPolicyService.ProtectionDescriptor::tableName).toList()); + String fromTable = resolveDefaultTableName(select, aliasToTable); + ConnectionChatAccessPolicyService.ProtectionDescriptor descriptor = lookupDescriptor(protectedObjects, fromTable); + if (descriptor != null && (descriptor.protectWholeTable() || !descriptor.restrictedColumns().isEmpty())) { + inspection.selectsWildcardFromProtectedTable = true; + inspection.reason = "SELECT * touches protected objects"; + inspection.protectedTables.add(descriptor.qualifiedTableName()); + } return; } if (expression instanceof AllTableColumns allTableColumns) { - String tableName = resolveTableName(allTableColumns.getTable(), aliasToTable); - ConnectionChatAccessPolicyService.ProtectionDescriptor descriptor = protectedObjects.get(normalizeName(tableName)); + String tableName = resolveQualifiedTableName(allTableColumns.getTable(), aliasToTable); + ConnectionChatAccessPolicyService.ProtectionDescriptor descriptor = lookupDescriptor(protectedObjects, tableName); if (descriptor != null) { inspection.selectsWildcardFromProtectedTable = true; inspection.reason = "SELECT table.* touches protected table"; - inspection.protectedTables.add(tableName); + inspection.protectedTables.add(descriptor.qualifiedTableName()); } return; } @@ -286,15 +416,15 @@ private QueryInspection inspectPlainSelect( collectColumns(expression, referencedColumns, aliasToTable, defaultTableName); boolean aggregateExpression = containsAggregate(expression); for (ColumnReference reference : referencedColumns) { - ConnectionChatAccessPolicyService.ProtectionDescriptor descriptor = protectedObjects.get(normalizeName(reference.tableName())); + ConnectionChatAccessPolicyService.ProtectionDescriptor descriptor = lookupDescriptor(protectedObjects, reference.tableName()); if (descriptor == null) { continue; } boolean protectedColumn = descriptor.protectWholeTable() || descriptor.restrictedColumns().stream().anyMatch(column -> normalizeName(column).equals(normalizeName(reference.columnName()))); if (protectedColumn) { - inspection.protectedTables.add(reference.tableName()); - inspection.protectedColumns.add(reference.tableName() + "." + reference.columnName()); + inspection.protectedTables.add(descriptor.qualifiedTableName()); + inspection.protectedColumns.add(descriptor.qualifiedTableName() + "." + reference.columnName()); if (!aggregateExpression) { inspection.rawProtectedColumnsSelected = true; inspection.reason = "Raw protected column selected"; @@ -311,12 +441,12 @@ private String resolveDefaultTableName(PlainSelect select, Map a } Set concreteTables = new LinkedHashSet<>(); if (select.getFromItem() instanceof Table table) { - concreteTables.add(resolveTableName(table, aliasToTable)); + concreteTables.add(resolveQualifiedTableName(table, aliasToTable)); } if (select.getJoins() != null) { for (Join join : select.getJoins()) { if (join.getRightItem() instanceof Table table) { - concreteTables.add(resolveTableName(table, aliasToTable)); + concreteTables.add(resolveQualifiedTableName(table, aliasToTable)); } } } @@ -336,9 +466,16 @@ private Map buildAliasMap(PlainSelect select) { private void extractAlias(FromItem fromItem, Map aliasMap) { if (fromItem instanceof Table table) { - aliasMap.put(normalizeName(table.getName()), table.getName()); + String qualified; + if (table.getSchemaName() != null && !table.getSchemaName().isBlank()) { + qualified = ConnectionChatAccessPolicyService.qualifyTable(table.getSchemaName(), table.getName()); + } else { + qualified = table.getName(); + } + aliasMap.put(normalizeName(table.getName()), qualified); + aliasMap.put(normalizeName(qualified), qualified); if (table.getAlias() != null) { - aliasMap.put(normalizeName(table.getAlias().getName()), table.getName()); + aliasMap.put(normalizeName(table.getAlias().getName()), qualified); } } else if (fromItem instanceof ParenthesedSelect parenthesedSelect && parenthesedSelect.getAlias() != null) { @@ -356,7 +493,7 @@ private void collectColumns( return; } if (expression instanceof Column column) { - String resolvedTableName = resolveTableName(column.getTable(), aliasToTable); + String resolvedTableName = resolveQualifiedTableName(column.getTable(), aliasToTable); if (resolvedTableName.isBlank()) { resolvedTableName = defaultTableName; } @@ -381,13 +518,38 @@ private void collectColumns( } } - private String resolveTableName(Table table, Map aliasToTable) { + private String resolveQualifiedTableName(Table table, Map aliasToTable) { if (table == null || table.getName() == null) { return ""; } + if (table.getSchemaName() != null && !table.getSchemaName().isBlank()) { + return ConnectionChatAccessPolicyService.qualifyTable(table.getSchemaName(), table.getName()); + } return aliasToTable.getOrDefault(normalizeName(table.getName()), table.getName()); } + private ConnectionChatAccessPolicyService.ProtectionDescriptor lookupDescriptor( + Map protectedObjects, + String tableRef + ) { + if (tableRef == null || tableRef.isBlank()) { + return null; + } + ConnectionChatAccessPolicyService.ProtectionDescriptor direct = protectedObjects.get(normalizeName(tableRef)); + if (direct != null) { + return direct; + } + if (!tableRef.contains(".")) { + List matches = protectedObjects.values().stream() + .filter(descriptor -> normalizeName(descriptor.tableName()).equals(normalizeName(tableRef))) + .toList(); + if (matches.size() == 1) { + return matches.getFirst(); + } + } + return null; + } + private boolean containsAggregate(Expression expression) { if (expression instanceof Function function) { String name = function.getName(); @@ -410,11 +572,9 @@ private boolean shouldRedactColumn(String column, Set protectedColumnNam if (protectedColumnNames.contains(normalized)) { return true; } - if (normalized.contains(".")) { - String bare = normalized.substring(normalized.indexOf('.') + 1); - return protectedColumnNames.contains(bare); - } - return false; + return protectedColumnNames.stream().anyMatch(protectedName -> + protectedName.endsWith("." + normalized) || protectedName.equals(normalized) + ); } private Object redactValue(Object value) { diff --git a/backend/src/main/java/com/dbaagent/service/security/AccessControlService.java b/backend/src/main/java/com/dbaagent/service/security/AccessControlService.java index ecefb6e..4a32265 100644 --- a/backend/src/main/java/com/dbaagent/service/security/AccessControlService.java +++ b/backend/src/main/java/com/dbaagent/service/security/AccessControlService.java @@ -7,6 +7,7 @@ import com.dbaagent.repository.AnalysisHistoryRepository; import com.dbaagent.repository.ChatFeedbackRepository; import com.dbaagent.repository.ChatRepository; +import com.dbaagent.security.ImpersonationContext; import lombok.RequiredArgsConstructor; import org.springframework.beans.factory.annotation.Value; import org.springframework.security.core.Authentication; @@ -68,7 +69,7 @@ public void assertCanManageConnectionConfig(String connectionId) { } public ConnectionAccessService.ResolvedConnectionAccess resolveCurrentUserAccess(String connectionId) { - if (!authEnabled) { + if (!authEnabled && !ImpersonationContext.isActive()) { try { return connectionAccessService.resolveAccess(connectionId, null, true); } catch (RuntimeException e) { @@ -181,6 +182,11 @@ public String requireCurrentUsername() { } public boolean isCurrentUserAdmin() { + if (ImpersonationContext.isActive()) { + return ImpersonationContext.current() + .map(state -> state.target() != null && state.target().isAdmin()) + .orElse(false); + } if (!authEnabled) { return true; } @@ -199,7 +205,7 @@ private Chat findAccessibleChat(String chatId) { } private Optional findAccessibleChatIfPresent(String chatId) { - if (!authEnabled) { + if (!authEnabled && !ImpersonationContext.isActive()) { return chatRepository.findById(chatId); } String username = requireCurrentUsername(); diff --git a/backend/src/test/java/com/dbaagent/service/ConnectionChatAccessPolicyServiceTest.java b/backend/src/test/java/com/dbaagent/service/ConnectionChatAccessPolicyServiceTest.java index 108598a..f1ee627 100644 --- a/backend/src/test/java/com/dbaagent/service/ConnectionChatAccessPolicyServiceTest.java +++ b/backend/src/test/java/com/dbaagent/service/ConnectionChatAccessPolicyServiceTest.java @@ -46,11 +46,11 @@ void setUp() throws SQLException { SchemaMetadata schema = new SchemaMetadata(); schema.setTables(List.of( - table("customer_profiles", "email", "phone_number", "full_name"), - table("payment_profiles", "credit_card_last4", "bank_account_masked"), - table("bookings", "booking_id", "status") + schemaTable(null, "customer_profiles", "email", "varchar", "phone_number", "varchar", "full_name", "varchar"), + schemaTable(null, "payment_profiles", "credit_card_last4", "varchar", "bank_account_masked", "varchar"), + schemaTable(null, "bookings", "booking_id", "varchar", "status", "varchar") )); - when(schemaScannerService.scanSchema("conn-1")).thenReturn(schema); + lenient().when(schemaScannerService.scanSchema("conn-1")).thenReturn(schema); lenient().when(tableClassificationRepository.findLatestByConnectionIdOrderByTableNameAsc(anyString())) .thenReturn(List.of( classification( @@ -86,6 +86,97 @@ void previewPolicy_normalizesPiiAndFinancialRestrictions() { assertThat(preview.isRedactMode()).isTrue(); } + @Test + void previewPolicy_scopesTypedColumnConstraintsToAllowedSchema() throws SQLException { + SchemaMetadata multiSchema = new SchemaMetadata(); + multiSchema.setTables(List.of( + schemaTable("crm", "customers", "amount", "numeric", "name", "varchar"), + schemaTable("sales", "orders", "amount", "numeric", "currency", "varchar"), + schemaTable("marts", "fct_enrollment", "amount", "numeric", "currency", "varchar"), + schemaTable("marts", "dim_ott_subscription", "amount", "numeric", "currency", "varchar") + )); + when(schemaScannerService.scanSchema("conn-ms")).thenReturn(multiSchema); + lenient().when(tableClassificationRepository.findLatestByConnectionIdOrderByTableNameAsc("conn-ms")) + .thenReturn(List.of()); + + String policyText = """ + This user should have access only to schema marts. In this table, the user cannot query \ + integer or float amount columns but can query columns that are string and represent currency code. \ + Strictly, The user cannot access any other schema other than marts + """; + + PolicyPreviewResponse preview = service.previewPolicy("conn-ms", policyText); + + assertThat(preview.getImpactedColumns()) + .contains("marts.fct_enrollment.amount", "marts.dim_ott_subscription.amount") + .doesNotContain( + "crm.customers.amount", + "sales.orders.amount", + "marts.fct_enrollment.currency", + "marts.dim_ott_subscription.currency" + ); + } + + @Test + void previewPolicy_appliesTypedColumnConstraintsToAnyColumnName() throws SQLException { + SchemaMetadata schema = new SchemaMetadata(); + schema.setTables(List.of( + schemaTable("hr", "employees", "salary", "numeric", "email", "varchar"), + schemaTable("finance", "ledger", "amount", "numeric", "account_code", "varchar") + )); + when(schemaScannerService.scanSchema("conn-hr")).thenReturn(schema); + lenient().when(tableClassificationRepository.findLatestByConnectionIdOrderByTableNameAsc("conn-hr")) + .thenReturn(List.of()); + + PolicyPreviewResponse preview = service.previewPolicy( + "conn-hr", + "This user should have access only to schema hr. The user cannot query numeric salary columns." + ); + + assertThat(preview.getImpactedColumns()) + .contains("hr.employees.salary") + .doesNotContain("finance.ledger.amount", "hr.employees.email"); + } + + @Test + void previewPolicy_typeOnlyConstraintIsStillSchemaScoped() throws SQLException { + SchemaMetadata schema = new SchemaMetadata(); + schema.setTables(List.of( + schemaTable("finance", "ledger", "amount", "numeric", "account_code", "varchar"), + schemaTable("sales", "orders", "amount", "numeric", "status", "varchar") + )); + when(schemaScannerService.scanSchema("conn-fin")).thenReturn(schema); + lenient().when(tableClassificationRepository.findLatestByConnectionIdOrderByTableNameAsc("conn-fin")) + .thenReturn(List.of()); + + PolicyPreviewResponse preview = service.previewPolicy( + "conn-fin", + "Access only to schema finance. Redact numeric columns." + ); + + assertThat(preview.getImpactedColumns()) + .contains("finance.ledger.amount") + .doesNotContain("sales.orders.amount", "finance.ledger.account_code"); + } + + @Test + void previewPolicy_collapsesWhitespaceInDenyClauses() throws SQLException { + SchemaMetadata schema = new SchemaMetadata(); + schema.setTables(List.of( + schemaTable("marts", "fct_enrollment", "amount", "numeric", "currency", "varchar") + )); + when(schemaScannerService.scanSchema("conn-ws")).thenReturn(schema); + lenient().when(tableClassificationRepository.findLatestByConnectionIdOrderByTableNameAsc("conn-ws")) + .thenReturn(List.of()); + + PolicyPreviewResponse preview = service.previewPolicy( + "conn-ws", + "Access only to schema marts. The user cannot query float amount columns." + ); + + assertThat(preview.getImpactedColumns()).contains("marts.fct_enrollment.amount"); + } + @Test void previewPolicy_resolvesExplicitTableAndColumnMentions() { PolicyPreviewResponse preview = service.previewPolicy( @@ -99,14 +190,17 @@ void previewPolicy_resolvesExplicitTableAndColumnMentions() { assertThat(preview.getImpactedColumns()).contains("customer_profiles.email"); } - private TableMetadata table(String name, String... columns) { + private TableMetadata schemaTable(String schema, String name, String... columnSpecs) { TableMetadata table = new TableMetadata(); + table.setSchema(schema); table.setName(name); - table.setColumns( - java.util.Arrays.stream(columns) - .map(column -> new ColumnMetadata(column, "varchar", null, true, false, null, 0)) - .toList() - ); + java.util.List columnMetadata = new java.util.ArrayList<>(); + for (int i = 0; i < columnSpecs.length; i += 2) { + String columnName = columnSpecs[i]; + String dataType = i + 1 < columnSpecs.length ? columnSpecs[i + 1] : "varchar"; + columnMetadata.add(new ColumnMetadata(columnName, dataType, null, true, false, null, 0)); + } + table.setColumns(columnMetadata); return table; } diff --git a/backend/src/test/java/com/dbaagent/service/ImpersonationServiceTest.java b/backend/src/test/java/com/dbaagent/service/ImpersonationServiceTest.java new file mode 100644 index 0000000..029ea4b --- /dev/null +++ b/backend/src/test/java/com/dbaagent/service/ImpersonationServiceTest.java @@ -0,0 +1,281 @@ +package com.dbaagent.service; + +import com.dbaagent.model.SecurityEventOutcome; +import com.dbaagent.model.SecurityEventType; +import com.dbaagent.model.User; +import com.dbaagent.model.UserAccountStatus; +import com.dbaagent.repository.UserRepository; +import com.dbaagent.security.CustomUserDetailsService; +import com.dbaagent.security.ImpersonationContext; +import jakarta.servlet.http.Cookie; +import org.junit.jupiter.api.AfterEach; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.mockito.ArgumentCaptor; +import org.mockito.InjectMocks; +import org.mockito.Mock; +import org.mockito.junit.jupiter.MockitoExtension; +import org.springframework.mock.web.MockHttpServletRequest; +import org.springframework.mock.web.MockHttpServletResponse; +import org.springframework.security.authentication.UsernamePasswordAuthenticationToken; +import org.springframework.security.core.authority.SimpleGrantedAuthority; +import org.springframework.security.core.context.SecurityContextHolder; +import org.springframework.security.core.userdetails.UserDetails; +import org.springframework.test.util.ReflectionTestUtils; +import org.springframework.web.server.ResponseStatusException; + +import java.util.List; +import java.util.Map; +import java.util.Optional; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertThrows; +import static org.junit.jupiter.api.Assertions.assertTrue; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.ArgumentMatchers.eq; +import static org.mockito.Mockito.never; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.when; + +@ExtendWith(MockitoExtension.class) +class ImpersonationServiceTest { + + @Mock + private UserRepository userRepository; + + @Mock + private CustomUserDetailsService userDetailsService; + + @Mock + private AuthSessionService authSessionService; + + @Mock + private SecurityEventService securityEventService; + + @InjectMocks + private ImpersonationService impersonationService; + + @BeforeEach + void setUp() { + ReflectionTestUtils.setField(impersonationService, "authEnabled", true); + ReflectionTestUtils.setField(impersonationService, "impersonateCookieName", "impersonate_user"); + } + + @AfterEach + void tearDown() { + SecurityContextHolder.clearContext(); + ImpersonationContext.clear(); + } + + @Test + void startWritesCookieAndAudits() { + User admin = user(1L, "admin", "ADMIN"); + User editor = user(2L, "marts-editor", "DEVELOPER"); + when(userRepository.findById(2L)).thenReturn(Optional.of(editor)); + + MockHttpServletRequest request = new MockHttpServletRequest(); + MockHttpServletResponse response = new MockHttpServletResponse(); + + ImpersonationContext.State state = impersonationService.start(admin, 2L, request, response); + + assertEquals("marts-editor", state.targetUsername()); + verify(authSessionService).writeImpersonationCookie(response, "impersonate_user", 2L); + ArgumentCaptor captor = + ArgumentCaptor.forClass(SecurityEventService.EventRequest.class); + verify(securityEventService).log(captor.capture()); + assertEquals(SecurityEventType.IMPERSONATION_STARTED, captor.getValue().eventType()); + assertEquals(SecurityEventOutcome.SUCCESS, captor.getValue().outcome()); + assertEquals(1L, captor.getValue().actorUserId()); + assertEquals(2L, captor.getValue().userId()); + } + + @Test + void startRejectsSelf() { + User admin = user(1L, "admin", "ADMIN"); + when(userRepository.findById(1L)).thenReturn(Optional.of(admin)); + + ResponseStatusException ex = assertThrows(ResponseStatusException.class, + () -> impersonationService.start(admin, 1L, new MockHttpServletRequest(), new MockHttpServletResponse())); + assertEquals(400, ex.getStatusCode().value()); + verify(authSessionService, never()).writeImpersonationCookie(any(), any(), eq(1L)); + } + + @Test + void startRejectsAnotherAdmin() { + User admin = user(1L, "admin", "ADMIN"); + User otherAdmin = user(3L, "ops-admin", "ADMIN"); + when(userRepository.findById(3L)).thenReturn(Optional.of(otherAdmin)); + + ResponseStatusException ex = assertThrows(ResponseStatusException.class, + () -> impersonationService.start(admin, 3L, new MockHttpServletRequest(), new MockHttpServletResponse())); + assertEquals(400, ex.getStatusCode().value()); + } + + @Test + void startRejectsLockedUser() { + User admin = user(1L, "admin", "ADMIN"); + User locked = user(4L, "locked-editor", "DEVELOPER"); + locked.setAccountStatus(UserAccountStatus.LOCKED.name()); + when(userRepository.findById(4L)).thenReturn(Optional.of(locked)); + + ResponseStatusException ex = assertThrows(ResponseStatusException.class, + () -> impersonationService.start(admin, 4L, new MockHttpServletRequest(), new MockHttpServletResponse())); + assertEquals(400, ex.getStatusCode().value()); + } + + @Test + void startRejectsNonAdminActor() { + User editor = user(2L, "marts-editor", "DEVELOPER"); + ResponseStatusException ex = assertThrows(ResponseStatusException.class, + () -> impersonationService.start(editor, 5L, new MockHttpServletRequest(), new MockHttpServletResponse())); + assertEquals(403, ex.getStatusCode().value()); + } + + @Test + void applySwapsPrincipalToTargetUser() { + User admin = user(1L, "admin", "ADMIN"); + User editor = user(2L, "marts-editor", "DEVELOPER"); + when(userRepository.findByUsername("admin")).thenReturn(Optional.of(admin)); + when(userRepository.findById(2L)).thenReturn(Optional.of(editor)); + UserDetails editorDetails = new org.springframework.security.core.userdetails.User( + "marts-editor", + "x", + List.of(new SimpleGrantedAuthority("ROLE_DEVELOPER"), new SimpleGrantedAuthority("USE_CHAT")) + ); + when(userDetailsService.loadUserByUsername("marts-editor")).thenReturn(editorDetails); + + SecurityContextHolder.getContext().setAuthentication( + new UsernamePasswordAuthenticationToken( + "admin", + null, + List.of(new SimpleGrantedAuthority("ROLE_ADMIN")) + ) + ); + + MockHttpServletRequest request = new MockHttpServletRequest("GET", "/api/schema/objects"); + request.setServletPath("/schema/objects"); + request.setCookies(new Cookie("impersonate_user", "2")); + + impersonationService.applyToRequest(request); + + assertEquals("marts-editor", SecurityContextHolder.getContext().getAuthentication().getName()); + assertTrue(ImpersonationContext.isActive()); + assertEquals("admin", ImpersonationContext.current().orElseThrow().impersonatorUsername()); + } + + @Test + void applySkipsImpersonationControlPlane() { + authenticateAdmin(); + MockHttpServletRequest request = new MockHttpServletRequest("POST", "/api/admin/impersonate"); + request.setServletPath("/admin/impersonate"); + request.setCookies(new Cookie("impersonate_user", "2")); + + impersonationService.applyToRequest(request); + + assertEquals("admin", SecurityContextHolder.getContext().getAuthentication().getName()); + assertFalse(ImpersonationContext.isActive()); + verify(userRepository, never()).findById(2L); + } + + @Test + void applySkipsLogoutAndRefresh() { + authenticateAdmin(); + MockHttpServletRequest logout = new MockHttpServletRequest("POST", "/api/auth/logout"); + logout.setServletPath("/auth/logout"); + logout.setCookies(new Cookie("impersonate_user", "2")); + impersonationService.applyToRequest(logout); + assertEquals("admin", SecurityContextHolder.getContext().getAuthentication().getName()); + + MockHttpServletRequest refresh = new MockHttpServletRequest("POST", "/api/auth/refresh"); + refresh.setServletPath("/auth/refresh"); + refresh.setCookies(new Cookie("impersonate_user", "2")); + impersonationService.applyToRequest(refresh); + assertEquals("admin", SecurityContextHolder.getContext().getAuthentication().getName()); + verify(userDetailsService, never()).loadUserByUsername(any()); + } + + @Test + void applySkipsMcpBearerTokens() { + authenticateAdmin(); + MockHttpServletRequest request = new MockHttpServletRequest("GET", "/api/connections"); + request.setServletPath("/connections"); + request.addHeader("Authorization", "Bearer dsql_mcp_abc.secret"); + request.setCookies(new Cookie("impersonate_user", "2")); + + impersonationService.applyToRequest(request); + + assertEquals("admin", SecurityContextHolder.getContext().getAuthentication().getName()); + assertFalse(ImpersonationContext.isActive()); + } + + @Test + void applyDoesNotSwapForNonAdminSession() { + SecurityContextHolder.getContext().setAuthentication( + new UsernamePasswordAuthenticationToken( + "marts-editor", + null, + List.of(new SimpleGrantedAuthority("ROLE_DEVELOPER")) + ) + ); + MockHttpServletRequest request = new MockHttpServletRequest("GET", "/api/schema/objects"); + request.setServletPath("/schema/objects"); + request.setCookies(new Cookie("impersonate_user", "9")); + + impersonationService.applyToRequest(request); + + assertEquals("marts-editor", SecurityContextHolder.getContext().getAuthentication().getName()); + verify(userRepository, never()).findById(9L); + } + + @Test + void decorateAuthPayloadUsesActiveContext() { + User admin = user(1L, "admin", "ADMIN"); + User editor = user(2L, "marts-editor", "DEVELOPER"); + ImpersonationContext.enter(new ImpersonationContext.State(admin, editor)); + + Map payload = new java.util.LinkedHashMap<>(); + payload.put("username", "marts-editor"); + impersonationService.decorateAuthPayload(new MockHttpServletRequest(), editor, payload); + + assertEquals(Boolean.TRUE, payload.get("impersonating")); + assertEquals("admin", payload.get("impersonatorUsername")); + } + + @Test + void listCandidatesExcludesAdminsSelfAndInactive() { + User admin = user(1L, "admin", "ADMIN"); + User editor = user(2L, "marts-editor", "DEVELOPER"); + User otherAdmin = user(3L, "ops", "ADMIN"); + User locked = user(4L, "locked", "DEVELOPER"); + locked.setAccountStatus(UserAccountStatus.LOCKED.name()); + when(userRepository.findAll()).thenReturn(List.of(admin, editor, otherAdmin, locked)); + + List> candidates = impersonationService.listCandidates(admin); + + assertEquals(1, candidates.size()); + assertEquals("marts-editor", candidates.get(0).get("username")); + } + + private void authenticateAdmin() { + SecurityContextHolder.getContext().setAuthentication( + new UsernamePasswordAuthenticationToken( + "admin", + null, + List.of(new SimpleGrantedAuthority("ROLE_ADMIN")) + ) + ); + } + + private static User user(Long id, String username, String role) { + User user = new User(); + user.setId(id); + user.setUsername(username); + user.setEmail(username + "@demo.local"); + user.setRole(role); + user.setAccountStatus(UserAccountStatus.ACTIVE.name()); + user.setPassword("hashed"); + return user; + } +} diff --git a/backend/src/test/java/com/dbaagent/service/UserDataAccessPolicyServiceTest.java b/backend/src/test/java/com/dbaagent/service/UserDataAccessPolicyServiceTest.java index 1c89136..c0ee0ff 100644 --- a/backend/src/test/java/com/dbaagent/service/UserDataAccessPolicyServiceTest.java +++ b/backend/src/test/java/com/dbaagent/service/UserDataAccessPolicyServiceTest.java @@ -31,11 +31,14 @@ class UserDataAccessPolicyServiceTest { @BeforeEach void setUp() { service = new UserDataAccessPolicyService(policyService, securityEventService); - lenient().when(policyService.buildProtectionDescriptors(anyString(), any(), any(), any())) - .thenReturn(Map.of( - "customer_profiles", - descriptor("customer_profiles", false, "email", "phone_number") - )); + lenient().when(policyService.buildProtectionDescriptors(any(ConnectionChatAccessPolicyService.EffectivePolicy.class))) + .thenAnswer(invocation -> { + ConnectionChatAccessPolicyService.EffectivePolicy policy = invocation.getArgument(0); + return Map.of( + "customer_profiles", + descriptor(null, "customer_profiles", false, "email", "phone_number") + ); + }); } @Test @@ -106,6 +109,156 @@ void decorateQuestionWithPolicy_injectsBoundariesForPlanner() { assertThat(decorated).contains("customer segments by month"); } + @Test + void enforcePreExecution_blocksQueriesOutsideAllowedSchema() { + ConnectionChatAccessPolicyService.EffectivePolicy schemaPolicy = new ConnectionChatAccessPolicyService.EffectivePolicy( + true, + "conn-1", + "analyst", + Set.of(), + Set.of(), + Set.of(), + Set.of("marts"), + true, + true, + "Only schema marts", + List.of(), + List.of() + ); + when(policyService.resolveEffectivePolicy("conn-1", "analyst", false)).thenReturn(schemaPolicy); + when(policyService.buildProtectionDescriptors(schemaPolicy)).thenReturn(Map.of()); + + UserDataAccessPolicyException exception = assertThrows( + UserDataAccessPolicyException.class, + () -> service.enforcePreExecution( + "conn-1", + new QueryRequest("SELECT name FROM crm.customers", null, null), + new QueryExecutionContext(QueryExecutionOrigin.CHAT, QueryExecutionContext.MutationMode.READ_ONLY_ONLY, "analyst", false, false) + ) + ); + + assertThat(exception.getErrorCode()).isEqualTo("POLICY_SCHEMA_BLOCKED"); + } + + @Test + void enforcePreExecution_allowsMartsQueriesWhenOtherSchemasHaveProtectedColumns() { + ConnectionChatAccessPolicyService.EffectivePolicy schemaPolicy = new ConnectionChatAccessPolicyService.EffectivePolicy( + true, + "conn-1", + "analyst", + Set.of(), + Set.of(), + Set.of("marts.fct_enrollment.amount"), + Set.of("marts"), + true, + true, + "Only marts; redact amount", + List.of(), + List.of("marts.fct_enrollment.amount") + ); + when(policyService.resolveEffectivePolicy("conn-1", "analyst", false)).thenReturn(schemaPolicy); + when(policyService.buildProtectionDescriptors(schemaPolicy)).thenReturn(Map.of( + "marts.fct_enrollment", + descriptor("marts", "fct_enrollment", false, "amount") + )); + + assertThat(service.enforcePreExecution( + "conn-1", + new QueryRequest("SELECT currency FROM marts.fct_enrollment", null, null), + new QueryExecutionContext(QueryExecutionOrigin.CHAT, QueryExecutionContext.MutationMode.READ_ONLY_ONLY, "analyst", false, false) + ).policy().allowedSchemas()).containsExactly("marts"); + } + + @Test + void filterDatabaseObjects_keepsOnlyAllowedSchemas() { + ConnectionChatAccessPolicyService.EffectivePolicy schemaPolicy = new ConnectionChatAccessPolicyService.EffectivePolicy( + true, + "conn-1", + "analyst", + Set.of(), + Set.of(), + Set.of(), + Set.of("marts"), + true, + true, + "Only schema marts", + List.of(), + List.of() + ); + when(policyService.resolveEffectivePolicy("conn-1", "analyst", false)).thenReturn(schemaPolicy); + + var marts = new com.dbaagent.model.DatabaseObject("fct_enrollment", "table", "marts", List.of(), 2L, null); + var crm = new com.dbaagent.model.DatabaseObject("customers", "table", "crm", List.of(), 2L, null); + var sales = new com.dbaagent.model.DatabaseObject("orders", "table", "sales", List.of(), 2L, null); + + assertThat(service.filterDatabaseObjects("conn-1", "analyst", false, List.of(marts, crm, sales))) + .extracting(com.dbaagent.model.DatabaseObject::getName) + .containsExactly("fct_enrollment"); + } + + @Test + void filterSchemaMetadata_dropsOutOfScopeTablesAndRelationships() { + ConnectionChatAccessPolicyService.EffectivePolicy schemaPolicy = new ConnectionChatAccessPolicyService.EffectivePolicy( + true, + "conn-1", + "analyst", + Set.of(), + Set.of(), + Set.of(), + Set.of("marts"), + true, + true, + "Only schema marts", + List.of(), + List.of() + ); + when(policyService.resolveEffectivePolicy("conn-1", "analyst", false)).thenReturn(schemaPolicy); + + var schema = new com.dbaagent.model.SchemaMetadata(); + var marts = new com.dbaagent.model.TableMetadata(); + marts.setSchema("marts"); + marts.setName("fct_enrollment"); + var crm = new com.dbaagent.model.TableMetadata(); + crm.setSchema("crm"); + crm.setName("customers"); + schema.setTables(List.of(marts, crm)); + var relationship = new com.dbaagent.model.RelationshipMetadata(); + relationship.setFromTable("crm.customers"); + relationship.setToTable("marts.fct_enrollment"); + schema.setRelationships(List.of(relationship)); + + var filtered = service.filterSchemaMetadata("conn-1", "analyst", false, schema); + assertThat(filtered.getTables()).extracting(com.dbaagent.model.TableMetadata::getName) + .containsExactly("fct_enrollment"); + assertThat(filtered.getRelationships()).isEmpty(); + } + + @Test + void assertTableSchemaAllowed_blocksOutOfScopeTableMetadata() { + ConnectionChatAccessPolicyService.EffectivePolicy schemaPolicy = new ConnectionChatAccessPolicyService.EffectivePolicy( + true, + "conn-1", + "analyst", + Set.of(), + Set.of(), + Set.of(), + Set.of("marts"), + true, + true, + "Only schema marts", + List.of(), + List.of() + ); + when(policyService.resolveEffectivePolicy("conn-1", "analyst", false)).thenReturn(schemaPolicy); + + UserDataAccessPolicyException exception = assertThrows( + UserDataAccessPolicyException.class, + () -> service.assertTableSchemaAllowed("conn-1", "analyst", false, "crm.customers") + ); + assertThat(exception.getErrorCode()).isEqualTo("POLICY_SCHEMA_BLOCKED"); + service.assertTableSchemaAllowed("conn-1", "analyst", false, "marts.fct_enrollment"); + } + private ConnectionChatAccessPolicyService.EffectivePolicy policy() { return new ConnectionChatAccessPolicyService.EffectivePolicy( true, @@ -114,6 +267,7 @@ private ConnectionChatAccessPolicyService.EffectivePolicy policy() { Set.of("PII_MEDIUM"), Set.of(), Set.of("customer_profiles.email", "customer_profiles.phone_number"), + Set.of(), true, true, "No PII", @@ -122,12 +276,17 @@ private ConnectionChatAccessPolicyService.EffectivePolicy policy() { ); } - private ConnectionChatAccessPolicyService.ProtectionDescriptor descriptor(String tableName, boolean protectWholeTable, String... columns) { + private ConnectionChatAccessPolicyService.ProtectionDescriptor descriptor( + String schemaName, + String tableName, + boolean protectWholeTable, + String... columns + ) { try { var constructor = ConnectionChatAccessPolicyService.ProtectionDescriptor.class - .getDeclaredConstructor(String.class, boolean.class, Set.class); + .getDeclaredConstructor(String.class, String.class, boolean.class, Set.class); constructor.setAccessible(true); - return constructor.newInstance(tableName, protectWholeTable, new java.util.LinkedHashSet<>(List.of(columns))); + return constructor.newInstance(schemaName, tableName, protectWholeTable, new java.util.LinkedHashSet<>(List.of(columns))); } catch (ReflectiveOperationException e) { throw new RuntimeException(e); } diff --git a/backend/src/test/java/com/dbaagent/service/security/AccessControlServiceTest.java b/backend/src/test/java/com/dbaagent/service/security/AccessControlServiceTest.java index 4673cbb..1590e3a 100644 --- a/backend/src/test/java/com/dbaagent/service/security/AccessControlServiceTest.java +++ b/backend/src/test/java/com/dbaagent/service/security/AccessControlServiceTest.java @@ -53,6 +53,7 @@ void setUp() { @AfterEach void tearDown() { SecurityContextHolder.clearContext(); + com.dbaagent.security.ImpersonationContext.clear(); } @Test @@ -97,6 +98,39 @@ void adminCanAccessAnyConnection() { assertDoesNotThrow(() -> accessControlService.assertCanAccessConnection("conn-1")); } + /** + * Profile switch has to punch through the auth-disabled admin bypass. + * Otherwise an admin "viewing as" an editor still sees every connection. + */ + @Test + void impersonationDisablesAdminBypassWhileAuthIsOff() { + ReflectionTestUtils.setField(accessControlService, "authEnabled", false); + + com.dbaagent.model.User impersonator = new com.dbaagent.model.User(); + impersonator.setId(1L); + impersonator.setUsername("admin"); + impersonator.setRole("ADMIN"); + com.dbaagent.model.User target = new com.dbaagent.model.User(); + target.setId(2L); + target.setUsername("marts-editor"); + target.setRole("DEVELOPER"); + com.dbaagent.security.ImpersonationContext.enter( + new com.dbaagent.security.ImpersonationContext.State(impersonator, target) + ); + SecurityContextHolder.getContext().setAuthentication( + new UsernamePasswordAuthenticationToken("marts-editor", null, List.of()) + ); + + when(connectionAccessService.resolveAccess("conn-1", "marts-editor", false)) + .thenReturn(resolved("conn-1", EffectiveConnectionAccess.CHAT_EDITOR, ConnectionOwnershipType.ASSIGNED)); + + assertFalse(accessControlService.isCurrentUserAdmin()); + assertEquals("marts-editor", accessControlService.requireCurrentUsername()); + assertDoesNotThrow(() -> accessControlService.assertCanUseChatEditor("conn-1")); + verify(connectionAccessService).resolveAccess("conn-1", "marts-editor", false); + verify(connectionAccessService, never()).resolveAccess(eq("conn-1"), eq(null), eq(true)); + } + /** * The dev-mode bypass has to be coherent. Every other check here honours * security.auth.enabled, so this one throwing 403 meant turning auth off turned chat diff --git a/docker-compose.yml b/docker-compose.yml index e64c3a4..c2bb83f 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -140,8 +140,10 @@ services: DEEPSQL_CHAT_API_KEY: ${DEEPSQL_CHAT_API_KEY:-} DEEPSQL_CHAT_ENDPOINT: ${DEEPSQL_CHAT_ENDPOINT:-} DEEPSQL_CHAT_MODEL: ${DEEPSQL_CHAT_MODEL:-gpt-5.4} - # Reach the backend over the compose network (MCP tools + provisioner) - DEEPSQL_API_BASE_URL: http://backend:8080/api/ + # Reach the backend over the compose network (MCP tools + provisioner). + # Override to http://host.docker.internal:8080/api/ when the Java backend + # runs on the host (native `mvn spring-boot:run`) instead of Compose. + DEEPSQL_API_BASE_URL: ${DEEPSQL_API_BASE_URL:-http://backend:8080/api/} # Shared secret with backend AgentBridgeService AGENT_PROVISION_SECRET: ${AGENT_PROVISION_SECRET:-} # Origins allowed by the agent API CSRF check @@ -153,6 +155,8 @@ services: DEEPSQL_AGENT_TRUSTED_PROXY_CIDRS: ${DEEPSQL_AGENT_TRUSTED_PROXY_CIDRS:-10.0.0.0/8,172.16.0.0/12,192.168.0.0/16} HERMES_WEBUI_TRUSTED_PROXY_CIDRS: ${DEEPSQL_AGENT_TRUSTED_PROXY_CIDRS:-10.0.0.0/8,172.16.0.0/12,192.168.0.0/16} HERMES_WEBUI_TRUSTED_AUTH_HEADER: X-Remote-User + extra_hosts: + - "host.docker.internal:host-gateway" ports: - "127.0.0.1:${DEEPSQL_AGENT_PORT:-8787}:8787" - "127.0.0.1:${DEEPSQL_AGENT_PROVISIONER_PORT:-8788}:8788" diff --git a/docker/postgres/init/11_create_acme_erp.sql b/docker/postgres/init/11_create_acme_erp.sql new file mode 100644 index 0000000..9f05639 --- /dev/null +++ b/docker/postgres/init/11_create_acme_erp.sql @@ -0,0 +1,142 @@ +-- ACME ERP — multi-schema fixture for Brain, MCP, and chat-access-policy tests. +-- Schemas: crm, sales, finance, inventory, hr, marts (analytics mart). +-- Duplicate bare table names (customers, orders) across crm/sales on purpose. + +SELECT 'Creating acme_erp multi-schema database' AS status; + +DROP DATABASE IF EXISTS acme_erp; +CREATE DATABASE acme_erp; + +\connect acme_erp + +CREATE EXTENSION IF NOT EXISTS pg_stat_statements; + +CREATE SCHEMA crm; +CREATE SCHEMA sales; +CREATE SCHEMA finance; +CREATE SCHEMA inventory; +CREATE SCHEMA hr; +CREATE SCHEMA marts; + +-- CRM +CREATE TABLE crm.customers ( + id SERIAL PRIMARY KEY, + name TEXT NOT NULL, + email TEXT, + amount NUMERIC(12, 2) DEFAULT 0, + created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP +); + +CREATE TABLE crm.accounts ( + id SERIAL PRIMARY KEY, + customer_id INT REFERENCES crm.customers(id), + balance NUMERIC(14, 2) NOT NULL DEFAULT 0 +); + +-- Sales (same bare names as crm) +CREATE TABLE sales.customers ( + id SERIAL PRIMARY KEY, + name TEXT NOT NULL, + email TEXT, + revenue NUMERIC(12, 2) DEFAULT 0 +); + +CREATE TABLE sales.orders ( + id SERIAL PRIMARY KEY, + customer_id INT REFERENCES sales.customers(id), + amount NUMERIC(12, 2) NOT NULL, + currency VARCHAR(3) NOT NULL DEFAULT 'USD', + status TEXT NOT NULL DEFAULT 'open', + ordered_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP +); + +-- Finance +CREATE TABLE finance.ledger ( + id SERIAL PRIMARY KEY, + account_code TEXT NOT NULL, + amount NUMERIC(14, 2) NOT NULL, + currency VARCHAR(3) NOT NULL DEFAULT 'USD', + posted_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP +); + +-- Inventory +CREATE TABLE inventory.products ( + id SERIAL PRIMARY KEY, + sku TEXT UNIQUE NOT NULL, + name TEXT NOT NULL, + stock_qty INT NOT NULL DEFAULT 0, + unit_cost NUMERIC(10, 2) +); + +-- HR +CREATE TABLE hr.employees ( + id SERIAL PRIMARY KEY, + full_name TEXT NOT NULL, + email TEXT, + salary NUMERIC(12, 2), + department TEXT +); + +-- Marts — policy-test schema (numeric amount + string currency, like production marts) +CREATE TABLE marts.fct_enrollment ( + id SERIAL PRIMARY KEY, + program_code TEXT NOT NULL, + amount NUMERIC(12, 2) NOT NULL, + currency VARCHAR(3) NOT NULL, + enrolled_at DATE NOT NULL +); + +CREATE TABLE marts.dim_ott_subscription ( + id SERIAL PRIMARY KEY, + subscriber_id TEXT NOT NULL, + amount NUMERIC(10, 2) NOT NULL, + currency VARCHAR(3) NOT NULL, + plan_name TEXT +); + +CREATE TABLE marts.rpt_revenue_monthly ( + month DATE NOT NULL, + region TEXT NOT NULL, + amount NUMERIC(14, 2) NOT NULL, + currency VARCHAR(3) NOT NULL, + PRIMARY KEY (month, region) +); + +-- Seed rows +INSERT INTO crm.customers (name, email, amount) VALUES + ('Acme Corp', 'acme@example.com', 1200.50), + ('Globex', 'globex@example.com', 800.00); + +INSERT INTO crm.accounts (customer_id, balance) VALUES (1, 500.00), (2, 250.75); + +INSERT INTO sales.customers (name, email, revenue) VALUES + ('Retail One', 'r1@example.com', 4200.00), + ('Retail Two', 'r2@example.com', 3100.25); + +INSERT INTO sales.orders (customer_id, amount, currency, status) VALUES + (1, 199.99, 'USD', 'shipped'), + (2, 89.50, 'EUR', 'open'); + +INSERT INTO finance.ledger (account_code, amount, currency) VALUES + ('CASH', 10000.00, 'USD'), + ('AR', 2500.00, 'USD'); + +INSERT INTO inventory.products (sku, name, stock_qty, unit_cost) VALUES + ('SKU-001', 'Widget', 120, 4.50), + ('SKU-002', 'Gadget', 45, 12.00); + +INSERT INTO hr.employees (full_name, email, salary, department) VALUES + ('Jane Doe', 'jane@acme.com', 95000.00, 'Engineering'), + ('John Smith', 'john@acme.com', 82000.00, 'Sales'); + +INSERT INTO marts.fct_enrollment (program_code, amount, currency, enrolled_at) VALUES + ('IEO', 150.00, 'USD', CURRENT_DATE - 30), + ('OTT', 9.99, 'INR', CURRENT_DATE - 7); + +INSERT INTO marts.dim_ott_subscription (subscriber_id, amount, currency, plan_name) VALUES + ('sub-100', 12.99, 'USD', 'Premium'), + ('sub-200', 499.00, 'INR', 'Annual'); + +INSERT INTO marts.rpt_revenue_monthly (month, region, amount, currency) VALUES + (DATE_TRUNC('month', CURRENT_DATE)::date, 'NA', 125000.00, 'USD'), + (DATE_TRUNC('month', CURRENT_DATE)::date, 'APAC', 98000.00, 'INR'); diff --git a/docs/root/CLAUDE.md b/docs/root/CLAUDE.md index 9c21e97..fcc4c2f 100644 --- a/docs/root/CLAUDE.md +++ b/docs/root/CLAUDE.md @@ -821,7 +821,10 @@ though the properties themselves still sit in `application*.properties`. - `PUT /api/admin/users/{id}/role` - Update user role (ADMIN only) - `DELETE /api/admin/users/{id}` - Delete user (ADMIN only) - `GET /api/admin/roles` - Get all roles with permissions (ADMIN only) - - `GET /api/auth/me` - Get current user's profile including role/permissions + - `GET /api/admin/impersonate` - List switchable users and current profile-switch status (ADMIN only) + - `POST /api/admin/impersonate` - `{ userId }` start viewing the product as that user (ADMIN only; cannot target admins or self) + - `DELETE /api/admin/impersonate` - Stop profile switch and restore the admin session + - `GET /api/auth/me` - Get current user's profile including role/permissions; while switching, this is the **target** user plus `impersonating` / `impersonatorUsername` - **Frontend Components**: - `PermissionGuard.jsx` - Wrapper component for permission-based rendering - `UsersTab.jsx` - Admin user management tab in Workspace diff --git a/scripts/remap-compose-hosts-for-native.sh b/scripts/remap-compose-hosts-for-native.sh new file mode 100644 index 0000000..1d551ba --- /dev/null +++ b/scripts/remap-compose-hosts-for-native.sh @@ -0,0 +1,22 @@ +# Sourced by scripts/start-backend.sh (outer shell and the inner bash -lc). +# Compose service hostnames only resolve on the compose network. Native +# `mvn spring-boot:run` still sources a Compose-oriented .env, so rewrite +# those hosts to loopback when they don't resolve. +if ! getent hosts postgres >/dev/null 2>&1; then + if [ -n "${DB_URL:-}" ]; then + export DB_URL="${DB_URL//:\/\/postgres:/:\/\/127.0.0.1:}" + fi +fi +if ! getent hosts valkey >/dev/null 2>&1; then + case "${SPRING_DATA_REDIS_HOST:-}" in + valkey|"") export SPRING_DATA_REDIS_HOST=127.0.0.1 ;; + esac +fi +if ! getent hosts deepsql-agent >/dev/null 2>&1; then + if [ -n "${AGENT_WEBUI_URL:-}" ]; then + export AGENT_WEBUI_URL="${AGENT_WEBUI_URL//deepsql-agent/127.0.0.1}" + fi + if [ -n "${AGENT_PROVISIONER_URL:-}" ]; then + export AGENT_PROVISIONER_URL="${AGENT_PROVISIONER_URL//deepsql-agent/127.0.0.1}" + fi +fi diff --git a/scripts/seed-acme-erp.sh b/scripts/seed-acme-erp.sh new file mode 100644 index 0000000..ebea846 --- /dev/null +++ b/scripts/seed-acme-erp.sh @@ -0,0 +1,96 @@ +#!/usr/bin/env bash +# seed-acme-erp.sh — Create the acme_erp multi-schema Postgres DB and register a DeepSQL connection. +# +# Native (non-Docker) cloud VM usage: +# sudo -u postgres psql -f docker/postgres/init/11_create_acme_erp.sql +# bash scripts/seed-acme-erp.sh +# +# Requires backend on :8080 with dev auth bypass (SECURITY_AUTH_ENABLED=false). + +set -euo pipefail + +SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)" +ROOT_DIR="$(cd "$SCRIPT_DIR/.." && pwd)" +ENV_FILE="${DEEPSQL_ENV_FILE:-$ROOT_DIR/.env}" +SQL_FILE="$ROOT_DIR/docker/postgres/init/11_create_acme_erp.sql" +CONNECTION_NAME="${DEEPSQL_ACME_CONNECTION_NAME:-ACME ERP (Multi-Schema)}" + +if [[ -f "$ENV_FILE" ]]; then + set -a + # shellcheck disable=SC1090 + source "$ENV_FILE" + set +a +fi + +: "${DEEPSQL_BACKEND_PORT:=8080}" +: "${DB_PASSWORD:=postgres}" + +base="http://127.0.0.1:${DEEPSQL_BACKEND_PORT}/api" + +echo "==========================================" +echo "ACME ERP multi-schema fixture" +echo "==========================================" + +if [[ ! -f "$SQL_FILE" ]]; then + echo "Missing SQL file: $SQL_FILE" >&2 + exit 1 +fi + +echo "" +echo "Step 1: Applying $SQL_FILE ..." +if command -v pg_ctlcluster >/dev/null 2>&1; then + sudo pg_ctlcluster 16 main status >/dev/null 2>&1 || sudo pg_ctlcluster 16 main start +fi +sudo -u postgres psql -v ON_ERROR_STOP=1 -f "$SQL_FILE" + +echo "" +echo "Step 2: Registering DeepSQL connection '$CONNECTION_NAME' ..." + +existing="$(curl -fsS "$base/connections" 2>/dev/null || echo '[]')" +connection_id="$(printf '%s' "$existing" | python3 -c " +import json, sys, os +name = os.environ.get('CONNECTION_NAME', '') +try: + data = json.load(sys.stdin) +except Exception: + data = [] +items = data if isinstance(data, list) else data.get('connections') or data.get('data') or [] +for conn in items: + if conn.get('connectionName') == name: + print(conn.get('id') or conn.get('connectionId') or '') + break +" CONNECTION_NAME="$CONNECTION_NAME")" + +if [[ -n "$connection_id" ]]; then + echo " Connection already exists: $connection_id" +else + payload="$(cat </dev/null || true)" + if [[ -z "$connection_id" ]]; then + echo " Warning: could not create connection (HTTP ${http_code:-?}). Body: $body" >&2 + else + echo " Created connection: $connection_id" + fi +fi + +echo "" +echo "Schemas in acme_erp:" +sudo -u postgres psql -d acme_erp -At -c "SELECT nspname FROM pg_namespace WHERE nspname NOT LIKE 'pg_%' AND nspname <> 'information_schema' ORDER BY 1;" + +echo "" +echo "Done. Use connection '$CONNECTION_NAME'${connection_id:+ (id=$connection_id)} for multi-schema / policy tests." diff --git a/scripts/start-backend.sh b/scripts/start-backend.sh index c67ffa7..4c87d2c 100755 --- a/scripts/start-backend.sh +++ b/scripts/start-backend.sh @@ -9,6 +9,8 @@ ENV_FILE="$PROJECT_ROOT/.env" echo "Starting DBA Agent Backend..." echo "================================" +REMAP_SCRIPT="$SCRIPT_DIR/remap-compose-hosts-for-native.sh" + if [ -f "$ENV_FILE" ]; then echo "Loading environment from .env..." set -a @@ -18,13 +20,16 @@ if [ -f "$ENV_FILE" ]; then echo "Local source-run startup ignores SPRING_PROFILES_ACTIVE=prod from .env" unset SPRING_PROFILES_ACTIVE fi + # shellcheck source=remap-compose-hosts-for-native.sh + source "$REMAP_SCRIPT" + echo "Agent provisioner: ${AGENT_PROVISIONER_URL:-unset}" fi build_backend_launch_command() { local mvn_command="$1" local env_snippet="" if [ -f "$ENV_FILE" ]; then - env_snippet="set -a && source \"$ENV_FILE\" && set +a && if [ \"\${SPRING_PROFILES_ACTIVE:-}\" = \"prod\" ]; then unset SPRING_PROFILES_ACTIVE; fi && " + env_snippet="set -a && source \"$ENV_FILE\" && set +a && if [ \"\${SPRING_PROFILES_ACTIVE:-}\" = \"prod\" ]; then unset SPRING_PROFILES_ACTIVE; fi && source \"$REMAP_SCRIPT\" && " fi printf '%s' "${env_snippet}cd \"$PROJECT_ROOT/backend\" && exec ${mvn_command} spring-boot:run" } diff --git a/src/components/layout/AppSidebar.jsx b/src/components/layout/AppSidebar.jsx index b9b5c9e..d6d12ce 100644 --- a/src/components/layout/AppSidebar.jsx +++ b/src/components/layout/AppSidebar.jsx @@ -1,5 +1,5 @@ import { useState, useEffect, useRef } from 'react' -import { BookOpen, Brain, Code2, Database, Settings, PanelLeftClose, PanelLeftOpen, LogOut, User, ChevronDown, Check, Newspaper, Gauge, MessageSquare, LayoutDashboard } from 'lucide-react' +import { Brain, Code2, Database, Settings, PanelLeftClose, PanelLeftOpen, LogOut, User, ChevronDown, Check, Newspaper, Gauge, MessageSquare, LayoutDashboard } from 'lucide-react' import { useActiveSection, useSetActiveSection } from '@/lib/stores/useNavStore' import { useConnectionManager } from '@/lib/hooks/useConnectionManager' import { AGENTS_ENABLED, canAccessHomeSection, getConnectionAccessBadge, getConnectionAccessLabel } from '@/lib/features' @@ -16,7 +16,6 @@ const NAV_ITEMS = [ { id: 'company-knowledge', label: 'Brain', icon: Brain }, { id: 'performance', label: 'Performance', icon: Gauge }, { id: 'editor', label: 'Editor', icon: Code2 }, - { id: 'docs', label: 'Docs', icon: BookOpen }, ] export default function AppSidebar() { @@ -29,7 +28,7 @@ export default function AppSidebar() { const [showConnectionDropdown, setShowConnectionDropdown] = useState(false) const userMenuRef = useRef(null) const connectionDropdownRef = useRef(null) - const { logout, role, username, isAdmin } = useAuth() + const { logout, role, username, impersonating } = useAuth() const { connections, connectionId, selectedConnection, changeConnection, isLoading, refetch } = useConnectionManager() const visibleNavItems = NAV_ITEMS.filter(({ id }) => canAccessHomeSection(id, role, selectedConnection)) @@ -92,8 +91,7 @@ export default function AppSidebar() {