| 1 | import csv |
| 2 | import hashlib |
| 3 | import io |
| 4 | import json |
| 5 | from datetime import datetime |
| 6 | from typing import Any |
| 7 | from typing import Dict |
| 8 | from typing import List |
| 9 | from typing import Optional |
| 10 | |
| 11 | from fastapi import HTTPException |
| 12 | from loguru import logger |
| 13 | from sqlalchemy import desc |
| 14 | from sqlalchemy import select |
| 15 | from sqlalchemy.ext.asyncio import AsyncSession |
| 16 | |
| 17 | from app.agents.vulnerabilities.schema.vulnerabilities import ( |
| 18 | AgentVulnerabilitiesResponse, |
| 19 | ) |
| 20 | from app.agents.vulnerabilities.schema.vulnerabilities import AgentVulnerabilityOut |
| 21 | from app.agents.vulnerabilities.schema.vulnerabilities import ( |
| 22 | VulnerabilityReportGenerateRequest, |
| 23 | ) |
| 24 | from app.agents.vulnerabilities.schema.vulnerabilities import ( |
| 25 | VulnerabilityReportGenerateResponse, |
| 26 | ) |
| 27 | from app.agents.vulnerabilities.schema.vulnerabilities import ( |
| 28 | VulnerabilityReportListResponse, |
| 29 | ) |
| 30 | from app.agents.vulnerabilities.schema.vulnerabilities import ( |
| 31 | VulnerabilityReportResponse, |
| 32 | ) |
| 33 | from app.agents.vulnerabilities.schema.vulnerabilities import VulnerabilitySearchItem |
| 34 | from app.agents.vulnerabilities.schema.vulnerabilities import ( |
| 35 | VulnerabilitySearchResponse, |
| 36 | ) |
| 37 | from app.agents.vulnerabilities.schema.vulnerabilities import VulnerabilityStatsResponse |
| 38 | from app.agents.vulnerabilities.schema.vulnerabilities import VulnerabilitySyncResponse |
| 39 | from app.agents.vulnerabilities.schema.vulnerabilities import WazuhVulnerabilityData |
| 40 | from app.auth.models.users import User |
| 41 | from app.connectors.wazuh_indexer.utils.universal import collect_indices |
| 42 | from app.connectors.wazuh_indexer.utils.universal import ( |
| 43 | create_wazuh_indexer_client_async, |
| 44 | ) |
| 45 | from app.data_store.data_store_operations import store_file_in_minio |
| 46 | from app.db.universal_models import Agents |
| 47 | from app.db.universal_models import AgentVulnerabilities |
| 48 | from app.db.universal_models import Customers |
| 49 | from app.db.universal_models import VulnerabilityReport |
| 50 | from app.middleware.customer_access import customer_access_handler |
| 51 | from app.threat_intel.schema.epss import EpssThreatIntelRequest |
| 52 | from app.threat_intel.services.epss import collect_epss_score |
| 53 | |
| 54 | |
| 55 | async def get_epss_score_for_cve(cve_id: str) -> tuple[Optional[str], Optional[str]]: |
| 56 | """ |
| 57 | Get EPSS score and percentile for a CVE ID |
| 58 | |
| 59 | Args: |
| 60 | cve_id: CVE identifier to get EPSS score for |
| 61 | |
| 62 | Returns: |
| 63 | Tuple of (epss_score, epss_percentile) or (None, None) if not found |
| 64 | """ |
| 65 | try: |
| 66 | epss_request = EpssThreatIntelRequest(cve=cve_id) |
| 67 | epss_response = await collect_epss_score(epss_request) |
| 68 | |
| 69 | if epss_response.success and epss_response.data: |
| 70 | # Get the first (and usually only) result |
| 71 | epss_data = epss_response.data[0] |
| 72 | return epss_data.epss, epss_data.percentile |
| 73 | else: |
| 74 | logger.debug(f"No EPSS data found for CVE: {cve_id}") |
| 75 | return None, None |
| 76 | |
| 77 | except Exception as e: |
| 78 | logger.warning(f"Error fetching EPSS score for CVE {cve_id}: {e}") |
| 79 | return None, None |
| 80 | |
| 81 | |
| 82 | def process_wazuh_document(document: Dict[str, Any]) -> WazuhVulnerabilityData: |
| 83 | """ |
| 84 | Process a single Wazuh vulnerability document from Indexer |
| 85 | |
| 86 | Args: |
| 87 | document: Raw document from Wazuh Indexer index |
| 88 | |
| 89 | Returns: |
| 90 | WazuhVulnerabilityData: Processed vulnerability data |
| 91 | """ |
| 92 | logger.info(f"Processing vulnerability document ID: {document.get('_id', 'unknown')}") |
| 93 | try: |
| 94 | source = document.get("_source", {}) |
| 95 | vuln_data = source.get("vulnerability", {}) |
| 96 | package_data = source.get("package", {}) |
| 97 | score_data = vuln_data.get("score", {}) |
| 98 | |
| 99 | # Parse detected_at timestamp |
| 100 | detected_at_str = vuln_data.get("detected_at") |
| 101 | detected_at = datetime.fromisoformat(detected_at_str.replace("Z", "+00:00")) if detected_at_str else datetime.utcnow() |
| 102 | |
| 103 | # Parse published_at timestamp if available |
| 104 | published_at_str = vuln_data.get("published_at") |
| 105 | published_at = None |
| 106 | if published_at_str: |
| 107 | try: |
| 108 | published_at = datetime.fromisoformat(published_at_str.replace("Z", "+00:00")) |
| 109 | except ValueError: |
| 110 | logger.warning(f"Could not parse published_at: {published_at_str}") |
| 111 | |
| 112 | # Parse and limit references to first 5 items if comma-separated |
| 113 | references_raw = vuln_data.get("reference") |
| 114 | references = None |
| 115 | if references_raw: |
| 116 | if isinstance(references_raw, str) and "," in references_raw: |
| 117 | # Split by comma, take first 5 items, and rejoin |
| 118 | reference_list = [ref.strip() for ref in references_raw.split(",")] |
| 119 | references = ", ".join(reference_list[:5]) |
| 120 | else: |
| 121 | references = str(references_raw) |
| 122 | |
| 123 | # Also ensure the references field doesn't exceed database column limit (2048 chars) |
| 124 | if len(references) > 2048: |
| 125 | references = references[:2045] + "..." |
| 126 | |
| 127 | return WazuhVulnerabilityData( |
| 128 | cve_id=vuln_data.get("id", "UNKNOWN_CVE"), |
| 129 | severity=vuln_data.get("severity", "UNKNOWN"), |
| 130 | title=package_data.get("name", "Unknown Package"), |
| 131 | references=references, |
| 132 | detected_at=detected_at, |
| 133 | published_at=published_at, |
| 134 | base_score=score_data.get("base"), |
| 135 | package_name=package_data.get("name"), |
| 136 | package_version=package_data.get("version"), |
| 137 | package_architecture=package_data.get("architecture"), |
| 138 | ) |
| 139 | except Exception as e: |
| 140 | logger.error(f"Error processing vulnerability document: {e}") |
| 141 | logger.error(f"Document: {document}") |
| 142 | raise |
| 143 | |
| 144 | |
| 145 | async def get_vulnerabilities_indices() -> List[str]: |
| 146 | """Get all vulnerability indices from Wazuh Indexer""" |
| 147 | try: |
| 148 | indices = await collect_indices(all_indices=True) |
| 149 | vuln_indices = [index for index in indices.indices_list if index.startswith("wazuh-states-vulnerabilities")] |
| 150 | logger.info(f"Found {len(vuln_indices)} vulnerability indices") |
| 151 | return vuln_indices |
| 152 | except Exception as e: |
| 153 | logger.error(f"Error collecting vulnerability indices: {e}") |
| 154 | raise HTTPException(status_code=500, detail=f"Failed to collect vulnerability indices: {e}") |
| 155 | |
| 156 | |
| 157 | async def fetch_vulnerabilities_from_indexer( |
| 158 | agent_name: Optional[str] = None, |
| 159 | customer_code: Optional[str] = None, |
| 160 | severity_filter: Optional[List[str]] = None, |
| 161 | ) -> List[Dict[str, Any]]: |
| 162 | """ |
| 163 | Fetch vulnerabilities from Wazuh Indexer indices |
| 164 | |
| 165 | Args: |
| 166 | agent_name: Optional agent name filter |
| 167 | customer_code: Optional customer code filter (used for index filtering) |
| 168 | severity_filter: Optional list of severities to filter by |
| 169 | |
| 170 | Returns: |
| 171 | List of vulnerability documents |
| 172 | """ |
| 173 | es_client = None |
| 174 | try: |
| 175 | es_client = await create_wazuh_indexer_client_async("Wazuh-Indexer") |
| 176 | indices = await get_vulnerabilities_indices() |
| 177 | |
| 178 | if not indices: |
| 179 | logger.warning("No vulnerability indices found") |
| 180 | return [] |
| 181 | |
| 182 | vulnerabilities = [] |
| 183 | |
| 184 | # Build query |
| 185 | query = {"query": {"bool": {"must": []}}} |
| 186 | |
| 187 | if agent_name: |
| 188 | query["query"]["bool"]["must"].append({"match": {"agent.name": agent_name}}) |
| 189 | |
| 190 | if severity_filter: |
| 191 | query["query"]["bool"]["must"].append({"terms": {"vulnerability.severity": severity_filter}}) |
| 192 | |
| 193 | # If no filters, match all |
| 194 | if not query["query"]["bool"]["must"]: |
| 195 | query = {"query": {"match_all": {}}} |
| 196 | |
| 197 | # Search across all vulnerability indices |
| 198 | for index in indices: |
| 199 | try: |
| 200 | # Use scroll for large result sets |
| 201 | page = await es_client.search(index=index, body=query, scroll="2m", size=1000) |
| 202 | scroll_id = page["_scroll_id"] |
| 203 | scroll_size = len(page["hits"]["hits"]) |
| 204 | |
| 205 | vulnerabilities.extend(page["hits"]["hits"]) |
| 206 | |
| 207 | # Continue scrolling through results |
| 208 | while scroll_size > 0: |
| 209 | page = await es_client.scroll(scroll_id=scroll_id, scroll="2m") |
| 210 | scroll_id = page["_scroll_id"] |
| 211 | scroll_size = len(page["hits"]["hits"]) |
| 212 | vulnerabilities.extend(page["hits"]["hits"]) |
| 213 | |
| 214 | # Clear the scroll context when done with this index |
| 215 | try: |
| 216 | await es_client.clear_scroll(scroll_id=scroll_id) |
| 217 | except Exception as clear_error: |
| 218 | logger.warning(f"Could not clear scroll context: {clear_error}") |
| 219 | |
| 220 | except Exception as index_error: |
| 221 | logger.error(f"Error querying index {index}: {index_error}") |
| 222 | continue |
| 223 | |
| 224 | logger.info(f"Fetched {len(vulnerabilities)} vulnerabilities from Indexer") |
| 225 | return vulnerabilities |
| 226 | |
| 227 | except Exception as e: |
| 228 | logger.error(f"Error fetching vulnerabilities from Indexer: {e}") |
| 229 | raise HTTPException(status_code=500, detail=f"Failed to fetch vulnerabilities: {e}") |
| 230 | finally: |
| 231 | # Ensure the Elasticsearch client session is properly closed |
| 232 | if es_client: |
| 233 | try: |
| 234 | await es_client.close() |
| 235 | except Exception as close_error: |
| 236 | logger.warning(f"Error closing Elasticsearch client: {close_error}") |
| 237 | |
| 238 | |
| 239 | async def get_agent_by_name(db_session: AsyncSession, agent_name: str) -> Optional[Agents]: |
| 240 | """Get agent from database by hostname/name""" |
| 241 | try: |
| 242 | result = await db_session.execute(select(Agents).filter(Agents.hostname == agent_name)) |
| 243 | return result.scalars().first() |
| 244 | except Exception as e: |
| 245 | logger.error(f"Error fetching agent {agent_name}: {e}") |
| 246 | return None |
| 247 | |
| 248 | |
| 249 | async def _sync_vulnerabilities_bulk_mode( |
| 250 | db_session: AsyncSession, |
| 251 | agent_id: str, |
| 252 | agent_name: str, |
| 253 | customer_code: str, |
| 254 | vulnerability_docs: List[Dict[str, Any]], |
| 255 | ) -> "VulnerabilitySyncResponse": |
| 256 | """ |
| 257 | Ultra-fast bulk mode for processing large numbers of vulnerabilities. |
| 258 | Uses SQLAlchemy bulk operations for maximum performance. |
| 259 | """ |
| 260 | from app.agents.vulnerabilities.schema.vulnerabilities import ( |
| 261 | VulnerabilitySyncResponse, |
| 262 | ) |
| 263 | |
| 264 | try: |
| 265 | logger.info(f"BULK MODE: Processing {len(vulnerability_docs)} vulnerabilities for agent {agent_name}") |
| 266 | |
| 267 | # Process all documents first |
| 268 | processed_vulns = [] |
| 269 | errors = [] |
| 270 | |
| 271 | for doc in vulnerability_docs: |
| 272 | try: |
| 273 | vuln_data = process_wazuh_document(doc) |
| 274 | processed_vulns.append(vuln_data) |
| 275 | except Exception as doc_error: |
| 276 | error_msg = f"Error processing vulnerability {doc.get('_id', 'unknown')}: {doc_error}" |
| 277 | logger.error(error_msg) |
| 278 | errors.append(error_msg) |
| 279 | |
| 280 | if not processed_vulns: |
| 281 | return VulnerabilitySyncResponse( |
| 282 | success=True, |
| 283 | message=f"No valid vulnerabilities to process for agent {agent_name}", |
| 284 | synced_count=0, |
| 285 | errors=errors, |
| 286 | ) |
| 287 | |
| 288 | # Get existing vulnerabilities for comparison |
| 289 | existing_vulns_result = await db_session.execute(select(AgentVulnerabilities).filter(AgentVulnerabilities.agent_id == agent_id)) |
| 290 | existing_vulns = existing_vulns_result.scalars().all() |
| 291 | |
| 292 | # Create lookup for existing vulnerabilities |
| 293 | existing_lookup = {} |
| 294 | for vuln in existing_vulns: |
| 295 | key = f"{vuln.cve_id}_{vuln.package_name or 'None'}" |
| 296 | existing_lookup[key] = vuln |
| 297 | |
| 298 | # Prepare bulk operations |
| 299 | new_vulnerabilities = [] |
| 300 | update_data = [] |
| 301 | |
| 302 | for vuln_data in processed_vulns: |
| 303 | key = f"{vuln_data.cve_id}_{vuln_data.package_name or 'None'}" |
| 304 | |
| 305 | if key in existing_lookup: |
| 306 | # Prepare for bulk update |
| 307 | existing_vuln = existing_lookup[key] |
| 308 | update_data.append( |
| 309 | { |
| 310 | "id": existing_vuln.id, |
| 311 | "severity": vuln_data.severity, |
| 312 | "title": vuln_data.title, |
| 313 | "references": vuln_data.references, |
| 314 | "discovered_at": vuln_data.detected_at, |
| 315 | "epss_score": str(vuln_data.base_score) |
| 316 | if hasattr(vuln_data, "base_score") and vuln_data.base_score |
| 317 | else existing_vuln.epss_score, |
| 318 | "package_name": vuln_data.package_name, |
| 319 | }, |
| 320 | ) |
| 321 | else: |
| 322 | # Prepare for bulk insert |
| 323 | new_vuln = AgentVulnerabilities.create_from_model( |
| 324 | vulnerability_data=vuln_data, |
| 325 | agent_id=agent_id, |
| 326 | customer_code=customer_code, |
| 327 | ) |
| 328 | new_vulnerabilities.append(new_vuln) |
| 329 | |
| 330 | # Execute bulk operations |
| 331 | inserted_count = 0 |
| 332 | updated_count = 0 |
| 333 | |
| 334 | if new_vulnerabilities: |
| 335 | db_session.add_all(new_vulnerabilities) |
| 336 | inserted_count = len(new_vulnerabilities) |
| 337 | logger.info(f"BULK MODE: Prepared {inserted_count} new vulnerabilities for insertion") |
| 338 | |
| 339 | if update_data: |
| 340 | # Use bulk update for existing vulnerabilities |
| 341 | from sqlalchemy import update |
| 342 | |
| 343 | for data in update_data: |
| 344 | stmt = ( |
| 345 | update(AgentVulnerabilities) |
| 346 | .where(AgentVulnerabilities.id == data["id"]) |
| 347 | .values( |
| 348 | { |
| 349 | "severity": data["severity"], |
| 350 | "title": data["title"], |
| 351 | "references": data["references"], |
| 352 | "discovered_at": data["discovered_at"], |
| 353 | "epss_score": data["epss_score"], |
| 354 | "package_name": data["package_name"], |
| 355 | }, |
| 356 | ) |
| 357 | ) |
| 358 | await db_session.execute(stmt) |
| 359 | updated_count = len(update_data) |
| 360 | logger.info(f"BULK MODE: Executed {updated_count} vulnerability updates") |
| 361 | |
| 362 | # Single commit for all operations |
| 363 | await db_session.commit() |
| 364 | |
| 365 | total_synced = inserted_count + updated_count |
| 366 | logger.info(f"BULK MODE: Successfully synced {total_synced} vulnerabilities ({inserted_count} new, {updated_count} updated)") |
| 367 | |
| 368 | return VulnerabilitySyncResponse( |
| 369 | success=True, |
| 370 | message=f"BULK MODE: Successfully synced {total_synced} vulnerabilities for agent {agent_name} ({inserted_count} new, {updated_count} updated)", |
| 371 | synced_count=total_synced, |
| 372 | errors=errors, |
| 373 | ) |
| 374 | |
| 375 | except Exception as e: |
| 376 | await db_session.rollback() |
| 377 | logger.error(f"BULK MODE: Error syncing vulnerabilities for agent {agent_name}: {e}") |
| 378 | return VulnerabilitySyncResponse( |
| 379 | success=False, |
| 380 | message=f"BULK MODE: Failed to sync vulnerabilities for agent {agent_name}: {e}", |
| 381 | synced_count=0, |
| 382 | errors=[str(e)], |
| 383 | ) |
| 384 | |
| 385 | |
| 386 | async def sync_vulnerabilities_for_agent( |
| 387 | db_session: AsyncSession, |
| 388 | agent_name: str, |
| 389 | customer_code: Optional[str] = None, |
| 390 | batch_size: int = 100, |
| 391 | use_bulk_mode: bool = False, |
| 392 | ) -> VulnerabilitySyncResponse: |
| 393 | """ |
| 394 | Sync vulnerabilities for a specific agent |
| 395 | |
| 396 | Args: |
| 397 | db_session: Database session to use |
| 398 | agent_name: Name of the agent to sync vulnerabilities for |
| 399 | customer_code: Optional customer code override |
| 400 | batch_size: Number of vulnerabilities to process in each batch (default: 100) |
| 401 | use_bulk_mode: If True, use ultra-fast bulk operations (default: False) |
| 402 | |
| 403 | Returns: |
| 404 | VulnerabilitySyncResponse with sync results |
| 405 | """ |
| 406 | try: |
| 407 | # Get agent from database using the session |
| 408 | result = await db_session.execute(select(Agents).filter(Agents.hostname == agent_name)) |
| 409 | agent = result.scalars().first() |
| 410 | |
| 411 | if not agent: |
| 412 | return VulnerabilitySyncResponse( |
| 413 | success=False, |
| 414 | message=f"Agent {agent_name} not found in database", |
| 415 | synced_count=0, |
| 416 | errors=[f"Agent {agent_name} not found"], |
| 417 | ) |
| 418 | |
| 419 | # Use agent's customer code if not provided |
| 420 | if not customer_code: |
| 421 | customer_code = agent.customer_code |
| 422 | |
| 423 | # Cache agent values to prevent lazy loading issues in the loop |
| 424 | agent_id = agent.agent_id |
| 425 | |
| 426 | # Fetch vulnerabilities from Indexer |
| 427 | vulnerability_docs = await fetch_vulnerabilities_from_indexer(agent_name=agent_name) |
| 428 | |
| 429 | logger.info(f"Fetched {len(vulnerability_docs)} vulnerabilities for agent {agent_name}") |
| 430 | |
| 431 | if not vulnerability_docs: |
| 432 | return VulnerabilitySyncResponse( |
| 433 | success=True, |
| 434 | message=f"No vulnerabilities found for agent {agent_name}", |
| 435 | synced_count=0, |
| 436 | errors=[], |
| 437 | ) |
| 438 | |
| 439 | synced_count = 0 |
| 440 | errors = [] |
| 441 | |
| 442 | # Choose processing mode based on use_bulk_mode flag |
| 443 | if use_bulk_mode: |
| 444 | logger.info(f"Using BULK MODE for {len(vulnerability_docs)} vulnerabilities for agent {agent_name}") |
| 445 | return await _sync_vulnerabilities_bulk_mode(db_session, agent_id, agent_name, customer_code, vulnerability_docs) |
| 446 | |
| 447 | # OPTIMIZATION: Process vulnerabilities in batches for better performance |
| 448 | logger.info(f"Using BATCH MODE (batch_size={batch_size}) for {len(vulnerability_docs)} vulnerabilities for agent {agent_name}") |
| 449 | |
| 450 | # First, get all existing vulnerabilities for this agent to do bulk comparison |
| 451 | logger.info(f"Fetching existing vulnerabilities for agent {agent_name} for comparison") |
| 452 | existing_vulns_result = await db_session.execute(select(AgentVulnerabilities).filter(AgentVulnerabilities.agent_id == agent_id)) |
| 453 | existing_vulns = existing_vulns_result.scalars().all() |
| 454 | |
| 455 | # Create a lookup dictionary for fast comparison (agent_id + cve_id + package_name) |
| 456 | existing_vulns_lookup = {} |
| 457 | for vuln in existing_vulns: |
| 458 | key = f"{vuln.agent_id}_{vuln.cve_id}_{vuln.package_name or 'None'}" |
| 459 | existing_vulns_lookup[key] = vuln |
| 460 | |
| 461 | logger.info(f"Found {len(existing_vulns_lookup)} existing vulnerabilities for agent {agent_name}") |
| 462 | |
| 463 | # Process vulnerabilities in batches |
| 464 | for batch_start in range(0, len(vulnerability_docs), batch_size): |
| 465 | batch_end = min(batch_start + batch_size, len(vulnerability_docs)) |
| 466 | batch_docs = vulnerability_docs[batch_start:batch_end] |
| 467 | |
| 468 | logger.info( |
| 469 | f"Processing batch {batch_start // batch_size + 1}: vulnerabilities {batch_start + 1}-{batch_end} of {len(vulnerability_docs)} for agent {agent_name}", |
| 470 | ) |
| 471 | |
| 472 | try: |
| 473 | batch_updates = [] |
| 474 | batch_inserts = [] |
| 475 | batch_errors = [] |
| 476 | |
| 477 | # Process each document in the batch |
| 478 | for i, doc in enumerate(batch_docs): |
| 479 | try: |
| 480 | # Process the vulnerability document |
| 481 | vuln_data = process_wazuh_document(doc) |
| 482 | |
| 483 | # Create lookup key |
| 484 | lookup_key = f"{agent_id}_{vuln_data.cve_id}_{vuln_data.package_name or 'None'}" |
| 485 | |
| 486 | if lookup_key in existing_vulns_lookup: |
| 487 | # Update existing vulnerability |
| 488 | existing_vuln = existing_vulns_lookup[lookup_key] |
| 489 | existing_vuln.severity = vuln_data.severity |
| 490 | existing_vuln.title = vuln_data.title |
| 491 | existing_vuln.references = vuln_data.references |
| 492 | existing_vuln.discovered_at = vuln_data.detected_at |
| 493 | if hasattr(vuln_data, "base_score") and vuln_data.base_score: |
| 494 | existing_vuln.epss_score = str(vuln_data.base_score) |
| 495 | if hasattr(vuln_data, "package_name"): |
| 496 | existing_vuln.package_name = vuln_data.package_name |
| 497 | |
| 498 | db_session.add(existing_vuln) |
| 499 | batch_updates.append(vuln_data.cve_id) |
| 500 | else: |
| 501 | # Create new vulnerability record |
| 502 | new_vuln = AgentVulnerabilities.create_from_model( |
| 503 | vulnerability_data=vuln_data, |
| 504 | agent_id=agent_id, |
| 505 | customer_code=customer_code, |
| 506 | ) |
| 507 | db_session.add(new_vuln) |
| 508 | batch_inserts.append(vuln_data.cve_id) |
| 509 | |
| 510 | # Add to lookup to avoid duplicates within the same batch |
| 511 | existing_vulns_lookup[lookup_key] = new_vuln |
| 512 | |
| 513 | except Exception as doc_error: |
| 514 | error_msg = f"Error processing vulnerability {doc.get('_id', 'unknown')}: {doc_error}" |
| 515 | logger.error(error_msg) |
| 516 | batch_errors.append(error_msg) |
| 517 | continue |
| 518 | |
| 519 | # Commit the entire batch at once |
| 520 | await db_session.commit() |
| 521 | |
| 522 | batch_synced = len(batch_updates) + len(batch_inserts) |
| 523 | synced_count += batch_synced |
| 524 | errors.extend(batch_errors) |
| 525 | |
| 526 | logger.info( |
| 527 | f"Batch {batch_start // batch_size + 1} completed: {len(batch_updates)} updates, {len(batch_inserts)} inserts, {len(batch_errors)} errors", |
| 528 | ) |
| 529 | |
| 530 | except Exception as batch_error: |
| 531 | await db_session.rollback() |
| 532 | error_msg = f"Error processing batch {batch_start}-{batch_end}: {batch_error}" |
| 533 | logger.error(error_msg) |
| 534 | errors.append(error_msg) |
| 535 | continue |
| 536 | |
| 537 | return VulnerabilitySyncResponse( |
| 538 | success=True, |
| 539 | message=f"Successfully synced {synced_count} vulnerabilities for agent {agent_name} ({len(errors)}", |
| 540 | synced_count=synced_count, |
| 541 | errors=errors, |
| 542 | ) |
| 543 | |
| 544 | except Exception as e: |
| 545 | await db_session.rollback() |
| 546 | logger.error(f"Error syncing vulnerabilities for agent {agent_name}: {e}") |
| 547 | return VulnerabilitySyncResponse( |
| 548 | success=False, |
| 549 | message=f"Failed to sync vulnerabilities for agent {agent_name}: {e}", |
| 550 | synced_count=0, |
| 551 | errors=[str(e)], |
| 552 | ) |
| 553 | |
| 554 | |
| 555 | async def sync_all_vulnerabilities( |
| 556 | db_session: AsyncSession, |
| 557 | customer_code: Optional[str] = None, |
| 558 | batch_size: int = 100, |
| 559 | use_bulk_mode: bool = False, |
| 560 | ) -> VulnerabilitySyncResponse: |
| 561 | """ |
| 562 | Sync vulnerabilities for all agents or agents of a specific customer with performance options |
| 563 | |
| 564 | Args: |
| 565 | db_session: Database session to use |
| 566 | customer_code: Optional customer code to filter agents by. |
| 567 | If None, syncs vulnerabilities for all agents in database. |
| 568 | batch_size: Number of vulnerabilities to process in each batch (default: 100) |
| 569 | use_bulk_mode: Use ultra-fast bulk operations for large datasets (default: False) |
| 570 | |
| 571 | Returns: |
| 572 | VulnerabilitySyncResponse with sync results |
| 573 | """ |
| 574 | try: |
| 575 | mode_info = "bulk mode" if use_bulk_mode else f"batch mode (size: {batch_size})" |
| 576 | logger.info(f"Starting bulk vulnerability sync for customer: {customer_code or 'all agents'} using {mode_info}") |
| 577 | |
| 578 | # Build query to get agents using the session |
| 579 | if customer_code: |
| 580 | query = select(Agents).filter(Agents.customer_code == customer_code) |
| 581 | else: |
| 582 | query = select(Agents) |
| 583 | |
| 584 | # Execute query using the session |
| 585 | result = await db_session.execute(query) |
| 586 | agents = result.scalars().all() |
| 587 | |
| 588 | if not agents: |
| 589 | message = "No agents found" + (f" for customer {customer_code}" if customer_code else " in database") |
| 590 | return VulnerabilitySyncResponse(success=True, message=message, synced_count=0, errors=[]) |
| 591 | |
| 592 | total_synced = 0 |
| 593 | all_errors = [] |
| 594 | |
| 595 | # Process each agent synchronously to maintain session consistency |
| 596 | for agent in agents: |
| 597 | # Cache agent values to prevent lazy loading issues |
| 598 | agent_hostname = agent.hostname |
| 599 | agent_customer_code = agent.customer_code |
| 600 | |
| 601 | if not agent_hostname: |
| 602 | continue |
| 603 | |
| 604 | try: |
| 605 | logger.info(f"Starting sync for agent: {agent_hostname} using {mode_info}") |
| 606 | result = await sync_vulnerabilities_for_agent( |
| 607 | db_session=db_session, |
| 608 | agent_name=agent_hostname, |
| 609 | customer_code=agent_customer_code, |
| 610 | batch_size=batch_size, |
| 611 | use_bulk_mode=use_bulk_mode, |
| 612 | ) |
| 613 | |
| 614 | total_synced += result.synced_count |
| 615 | all_errors.extend(result.errors) |
| 616 | |
| 617 | except Exception as agent_error: |
| 618 | error_msg = f"Error syncing agent {agent_hostname}: {agent_error}" |
| 619 | logger.error(error_msg) |
| 620 | all_errors.append(error_msg) |
| 621 | |
| 622 | success_message = f"Completed vulnerability sync for {len(agents)} agents using {mode_info}" |
| 623 | if customer_code: |
| 624 | success_message += f" (customer: {customer_code})" |
| 625 | else: |
| 626 | success_message += " (all agents in database)" |
| 627 | |
| 628 | return VulnerabilitySyncResponse(success=True, message=success_message, synced_count=total_synced, errors=all_errors) |
| 629 | |
| 630 | except Exception as e: |
| 631 | logger.error(f"Error in bulk vulnerability sync: {e}") |
| 632 | return VulnerabilitySyncResponse(success=False, message=f"Failed to sync vulnerabilities: {e}", synced_count=0, errors=[str(e)]) |
| 633 | |
| 634 | |
| 635 | async def get_vulnerabilities_by_agent( |
| 636 | db_session: AsyncSession, |
| 637 | agent_id: str, |
| 638 | severity_filter: Optional[List[str]] = None, |
| 639 | ) -> AgentVulnerabilitiesResponse: |
| 640 | """ |
| 641 | Get vulnerabilities for a specific agent from database |
| 642 | |
| 643 | Args: |
| 644 | db_session: Database session to use |
| 645 | agent_id: Agent ID to get vulnerabilities for |
| 646 | severity_filter: Optional list of severities to filter by |
| 647 | |
| 648 | Returns: |
| 649 | AgentVulnerabilitiesResponse with vulnerabilities |
| 650 | """ |
| 651 | try: |
| 652 | query = select(AgentVulnerabilities).filter(AgentVulnerabilities.agent_id == agent_id) |
| 653 | |
| 654 | if severity_filter: |
| 655 | query = query.filter(AgentVulnerabilities.severity.in_(severity_filter)) |
| 656 | |
| 657 | result = await db_session.execute(query) |
| 658 | vulnerabilities = result.scalars().all() |
| 659 | |
| 660 | vuln_list = [ |
| 661 | AgentVulnerabilityOut( |
| 662 | id=vuln.id, |
| 663 | cve_id=vuln.cve_id, |
| 664 | severity=vuln.severity, |
| 665 | title=vuln.title, |
| 666 | references=vuln.references, |
| 667 | status=vuln.status, |
| 668 | discovered_at=vuln.discovered_at, |
| 669 | remediated_at=vuln.remediated_at, |
| 670 | epss_score=vuln.epss_score, |
| 671 | epss_percentile=vuln.epss_percentile, |
| 672 | package_name=vuln.package_name, |
| 673 | agent_id=vuln.agent_id, |
| 674 | customer_code=vuln.customer_code, |
| 675 | ) |
| 676 | for vuln in vulnerabilities |
| 677 | ] |
| 678 | |
| 679 | return AgentVulnerabilitiesResponse( |
| 680 | vulnerabilities=vuln_list, |
| 681 | success=True, |
| 682 | message=f"Retrieved {len(vuln_list)} vulnerabilities for agent {agent_id}", |
| 683 | total_count=len(vuln_list), |
| 684 | ) |
| 685 | |
| 686 | except Exception as e: |
| 687 | logger.error(f"Error getting vulnerabilities for agent {agent_id}: {e}") |
| 688 | raise HTTPException(status_code=500, detail=f"Failed to get vulnerabilities for agent {agent_id}: {e}") |
| 689 | |
| 690 | |
| 691 | async def get_vulnerability_statistics(db_session: AsyncSession, customer_code: Optional[str] = None) -> VulnerabilityStatsResponse: |
| 692 | """ |
| 693 | Get vulnerability statistics |
| 694 | |
| 695 | Args: |
| 696 | db_session: Database session to use |
| 697 | customer_code: Optional customer code to filter by |
| 698 | |
| 699 | Returns: |
| 700 | VulnerabilityStatsResponse with statistics |
| 701 | """ |
| 702 | try: |
| 703 | query = select(AgentVulnerabilities) |
| 704 | if customer_code: |
| 705 | query = query.filter(AgentVulnerabilities.customer_code == customer_code) |
| 706 | |
| 707 | result = await db_session.execute(query) |
| 708 | vulnerabilities = result.scalars().all() |
| 709 | |
| 710 | # Calculate statistics |
| 711 | total = len(vulnerabilities) |
| 712 | critical = sum(1 for v in vulnerabilities if v.severity.lower() == "critical") |
| 713 | high = sum(1 for v in vulnerabilities if v.severity.lower() == "high") |
| 714 | medium = sum(1 for v in vulnerabilities if v.severity.lower() == "medium") |
| 715 | low = sum(1 for v in vulnerabilities if v.severity.lower() == "low") |
| 716 | |
| 717 | # Group by customer if no specific customer requested |
| 718 | by_customer = {} |
| 719 | if not customer_code: |
| 720 | for vuln in vulnerabilities: |
| 721 | if vuln.customer_code: |
| 722 | by_customer[vuln.customer_code] = by_customer.get(vuln.customer_code, 0) + 1 |
| 723 | |
| 724 | return VulnerabilityStatsResponse( |
| 725 | total_vulnerabilities=total, |
| 726 | critical_count=critical, |
| 727 | high_count=high, |
| 728 | medium_count=medium, |
| 729 | low_count=low, |
| 730 | by_customer=by_customer, |
| 731 | success=True, |
| 732 | message="Vulnerability statistics retrieved successfully", |
| 733 | ) |
| 734 | |
| 735 | except Exception as e: |
| 736 | logger.error(f"Error getting vulnerability statistics: {e}") |
| 737 | raise HTTPException(status_code=500, detail=f"Failed to get vulnerability statistics: {e}") |
| 738 | |
| 739 | |
| 740 | async def delete_vulnerabilities(db_session: AsyncSession, agent_name: Optional[str] = None, customer_code: Optional[str] = None): |
| 741 | """ |
| 742 | Delete vulnerabilities based on scope: |
| 743 | - If neither agent_name nor customer_code provided: Delete ALL vulnerabilities |
| 744 | - If agent_name provided: Delete vulnerabilities for that specific agent |
| 745 | - If customer_code provided: Delete vulnerabilities for all agents of that customer |
| 746 | |
| 747 | Args: |
| 748 | db_session: Database session to use |
| 749 | agent_name: Optional agent name to delete vulnerabilities for |
| 750 | customer_code: Optional customer code to delete vulnerabilities for |
| 751 | |
| 752 | Returns: |
| 753 | VulnerabilityDeleteResponse with deletion results |
| 754 | """ |
| 755 | from app.agents.vulnerabilities.schema.vulnerabilities import ( |
| 756 | VulnerabilityDeleteResponse, |
| 757 | ) |
| 758 | |
| 759 | try: |
| 760 | deleted_count = 0 |
| 761 | |
| 762 | if agent_name: |
| 763 | # Delete vulnerabilities for specific agent |
| 764 | logger.info(f"Deleting vulnerabilities for agent: {agent_name}") |
| 765 | |
| 766 | # First get the agent to validate it exists and get agent_id |
| 767 | agent_result = await db_session.execute(select(Agents).filter(Agents.hostname == agent_name)) |
| 768 | agent = agent_result.scalars().first() |
| 769 | |
| 770 | if not agent: |
| 771 | return VulnerabilityDeleteResponse( |
| 772 | success=False, |
| 773 | message=f"Agent {agent_name} not found in database", |
| 774 | deleted_count=0, |
| 775 | errors=[f"Agent {agent_name} not found"], |
| 776 | ) |
| 777 | |
| 778 | # Delete vulnerabilities for this agent |
| 779 | from sqlalchemy import delete |
| 780 | |
| 781 | delete_stmt = delete(AgentVulnerabilities).where(AgentVulnerabilities.agent_id == agent.agent_id) |
| 782 | result = await db_session.execute(delete_stmt) |
| 783 | deleted_count = result.rowcount |
| 784 | await db_session.commit() |
| 785 | |
| 786 | return VulnerabilityDeleteResponse( |
| 787 | success=True, |
| 788 | message=f"Successfully deleted {deleted_count} vulnerabilities for agent {agent_name}", |
| 789 | deleted_count=deleted_count, |
| 790 | errors=[], |
| 791 | ) |
| 792 | |
| 793 | elif customer_code: |
| 794 | # Delete vulnerabilities for all agents of specific customer |
| 795 | logger.info(f"Deleting vulnerabilities for customer: {customer_code}") |
| 796 | |
| 797 | from sqlalchemy import delete |
| 798 | |
| 799 | delete_stmt = delete(AgentVulnerabilities).where(AgentVulnerabilities.customer_code == customer_code) |
| 800 | result = await db_session.execute(delete_stmt) |
| 801 | deleted_count = result.rowcount |
| 802 | await db_session.commit() |
| 803 | |
| 804 | return VulnerabilityDeleteResponse( |
| 805 | success=True, |
| 806 | message=f"Successfully deleted {deleted_count} vulnerabilities for customer {customer_code}", |
| 807 | deleted_count=deleted_count, |
| 808 | errors=[], |
| 809 | ) |
| 810 | |
| 811 | else: |
| 812 | # Delete ALL vulnerabilities |
| 813 | logger.warning("Deleting ALL vulnerabilities from database") |
| 814 | |
| 815 | from sqlalchemy import delete |
| 816 | |
| 817 | delete_stmt = delete(AgentVulnerabilities) |
| 818 | result = await db_session.execute(delete_stmt) |
| 819 | deleted_count = result.rowcount |
| 820 | await db_session.commit() |
| 821 | |
| 822 | return VulnerabilityDeleteResponse( |
| 823 | success=True, |
| 824 | message=f"Successfully deleted ALL {deleted_count} vulnerabilities from database", |
| 825 | deleted_count=deleted_count, |
| 826 | errors=[], |
| 827 | ) |
| 828 | |
| 829 | except Exception as e: |
| 830 | await db_session.rollback() |
| 831 | logger.error(f"Error deleting vulnerabilities: {e}") |
| 832 | return VulnerabilityDeleteResponse( |
| 833 | success=False, |
| 834 | message=f"Failed to delete vulnerabilities: {e}", |
| 835 | deleted_count=0, |
| 836 | errors=[str(e)], |
| 837 | ) |
| 838 | |
| 839 | |
| 840 | async def search_vulnerabilities_from_indexer( |
| 841 | db_session: AsyncSession, |
| 842 | current_user: User, |
| 843 | customer_code: Optional[str] = None, |
| 844 | agent_name: Optional[str] = None, |
| 845 | severity: Optional[str] = None, |
| 846 | cve_id: Optional[str] = None, |
| 847 | package_name: Optional[str] = None, |
| 848 | page: int = 1, |
| 849 | page_size: int = 50, |
| 850 | include_epss: bool = True, |
| 851 | ) -> VulnerabilitySearchResponse: |
| 852 | """ |
| 853 | Search vulnerabilities directly from Wazuh indexer with filtering and pagination |
| 854 | |
| 855 | Args: |
| 856 | db_session: Database session for agent lookup |
| 857 | current_user: Current authenticated user for customer access filtering |
| 858 | customer_code: Optional customer code filter |
| 859 | agent_name: Optional agent hostname filter |
| 860 | severity: Optional severity filter |
| 861 | cve_id: Optional CVE ID filter |
| 862 | package_name: Optional package name filter |
| 863 | page: Page number for pagination |
| 864 | page_size: Number of results per page |
| 865 | include_epss: Whether to include EPSS scores (default: True, may impact performance) |
| 866 | |
| 867 | Returns: |
| 868 | VulnerabilitySearchResponse with paginated results filtered by user access |
| 869 | """ |
| 870 | logger.info( |
| 871 | f"Searching vulnerabilities with filters: customer_code={customer_code}, " |
| 872 | f"agent_name={agent_name}, severity={severity}, cve_id={cve_id}, " |
| 873 | f"package_name={package_name}, page={page}, page_size={page_size}, " |
| 874 | f"include_epss={include_epss}", |
| 875 | ) |
| 876 | |
| 877 | # Apply customer access filtering based on user permissions |
| 878 | accessible_customers = await customer_access_handler.get_user_accessible_customers(current_user, db_session) |
| 879 | logger.info(f"User {current_user.username} has access to customers: {accessible_customers}") |
| 880 | |
| 881 | # Override customer_code based on user permissions |
| 882 | if "*" not in accessible_customers: |
| 883 | if customer_code and customer_code not in accessible_customers: |
| 884 | return VulnerabilitySearchResponse( |
| 885 | vulnerabilities=[], |
| 886 | total_count=0, |
| 887 | critical_count=0, |
| 888 | high_count=0, |
| 889 | medium_count=0, |
| 890 | low_count=0, |
| 891 | page=page, |
| 892 | page_size=page_size, |
| 893 | total_pages=0, |
| 894 | has_next=False, |
| 895 | has_previous=False, |
| 896 | success=True, |
| 897 | message=f"Access denied to customer {customer_code}", |
| 898 | filters_applied={}, |
| 899 | ) |
| 900 | |
| 901 | # Build filters applied dict for response |
| 902 | filters_applied = {} |
| 903 | if customer_code: |
| 904 | filters_applied["customer_code"] = customer_code |
| 905 | if agent_name: |
| 906 | filters_applied["agent_name"] = agent_name |
| 907 | if severity: |
| 908 | filters_applied["severity"] = severity |
| 909 | if cve_id: |
| 910 | filters_applied["cve_id"] = cve_id |
| 911 | if package_name: |
| 912 | filters_applied["package_name"] = package_name |
| 913 | |
| 914 | es_client = None |
| 915 | try: |
| 916 | # Initialize Elasticsearch client |
| 917 | es_client = await create_wazuh_indexer_client_async("Wazuh-Indexer") |
| 918 | |
| 919 | # Get all agents for customer code mapping |
| 920 | all_agents_query = select(Agents) |
| 921 | all_agents_result = await db_session.execute(all_agents_query) |
| 922 | all_agents = all_agents_result.scalars().all() |
| 923 | |
| 924 | # Build complete agent hostname to customer code mapping |
| 925 | customer_agent_map = {} |
| 926 | for agent in all_agents: |
| 927 | if agent.hostname: |
| 928 | customer_agent_map[agent.hostname] = agent.customer_code |
| 929 | |
| 930 | # Get agent information for filtering (if filters are applied) |
| 931 | agent_hostnames = [] |
| 932 | |
| 933 | # Build base query for agents |
| 934 | query = select(Agents) |
| 935 | |
| 936 | # Apply user access restrictions first |
| 937 | if "*" not in accessible_customers: |
| 938 | query = query.filter(Agents.customer_code.in_(accessible_customers)) |
| 939 | |
| 940 | # Apply additional filters if specified |
| 941 | if customer_code: |
| 942 | query = query.filter(Agents.customer_code == customer_code) |
| 943 | if agent_name: |
| 944 | query = query.filter(Agents.hostname == agent_name) |
| 945 | |
| 946 | result = await db_session.execute(query) |
| 947 | agents = result.scalars().all() |
| 948 | |
| 949 | if not agents and (customer_code or agent_name or "*" not in accessible_customers): |
| 950 | return VulnerabilitySearchResponse( |
| 951 | vulnerabilities=[], |
| 952 | total_count=0, |
| 953 | critical_count=0, |
| 954 | high_count=0, |
| 955 | medium_count=0, |
| 956 | low_count=0, |
| 957 | page=page, |
| 958 | page_size=page_size, |
| 959 | total_pages=0, |
| 960 | has_next=False, |
| 961 | has_previous=False, |
| 962 | success=True, |
| 963 | message="No agents found matching the specified criteria or user access permissions", |
| 964 | filters_applied=filters_applied, |
| 965 | ) |
| 966 | |
| 967 | # Build list of agent hostnames for Elasticsearch filtering |
| 968 | if "*" not in accessible_customers or customer_code or agent_name: |
| 969 | for agent in agents: |
| 970 | if agent.hostname: |
| 971 | agent_hostnames.append(agent.hostname) |
| 972 | |
| 973 | # Get vulnerability indices |
| 974 | vuln_indices = await get_vulnerabilities_indices() |
| 975 | if not vuln_indices: |
| 976 | return VulnerabilitySearchResponse( |
| 977 | vulnerabilities=[], |
| 978 | total_count=0, |
| 979 | critical_count=0, |
| 980 | high_count=0, |
| 981 | medium_count=0, |
| 982 | low_count=0, |
| 983 | page=page, |
| 984 | page_size=page_size, |
| 985 | total_pages=0, |
| 986 | has_next=False, |
| 987 | has_previous=False, |
| 988 | success=True, |
| 989 | message="No vulnerability indices found", |
| 990 | filters_applied=filters_applied, |
| 991 | ) |
| 992 | |
| 993 | # Build Elasticsearch query |
| 994 | es_query = {"bool": {"must": []}} |
| 995 | |
| 996 | # Add agent filter if specified |
| 997 | if agent_hostnames: |
| 998 | es_query["bool"]["must"].append({"terms": {"agent.name": agent_hostnames}}) |
| 999 | |
| 1000 | # Add severity filter |
| 1001 | if severity: |
| 1002 | es_query["bool"]["must"].append({"term": {"vulnerability.severity": severity}}) |
| 1003 | |
| 1004 | # Add CVE ID filter |
| 1005 | if cve_id: |
| 1006 | es_query["bool"]["must"].append({"term": {"vulnerability.id": cve_id}}) |
| 1007 | |
| 1008 | # Add package name filter |
| 1009 | if package_name: |
| 1010 | es_query["bool"]["must"].append({"wildcard": {"package.name": f"*{package_name}*"}}) |
| 1011 | |
| 1012 | # First, get total count and severity aggregations |
| 1013 | count_response = await es_client.count(index=",".join(vuln_indices), body={"query": es_query}) |
| 1014 | total_count = count_response["count"] |
| 1015 | |
| 1016 | # Get severity aggregations |
| 1017 | agg_response = await es_client.search( |
| 1018 | index=",".join(vuln_indices), |
| 1019 | body={ |
| 1020 | "query": es_query, |
| 1021 | "size": 0, |
| 1022 | "aggs": {"severity_counts": {"terms": {"field": "vulnerability.severity", "size": 10}}}, |
| 1023 | }, |
| 1024 | ) |
| 1025 | |
| 1026 | # Extract severity counts |
| 1027 | severity_counts = {"Critical": 0, "High": 0, "Medium": 0, "Low": 0} |
| 1028 | if "aggregations" in agg_response and "severity_counts" in agg_response["aggregations"]: |
| 1029 | for bucket in agg_response["aggregations"]["severity_counts"]["buckets"]: |
| 1030 | severity_value = bucket["key"] |
| 1031 | count = bucket["doc_count"] |
| 1032 | if severity_value in severity_counts: |
| 1033 | severity_counts[severity_value] = count |
| 1034 | |
| 1035 | # Calculate pagination info |
| 1036 | total_pages = (total_count + page_size - 1) // page_size |
| 1037 | has_next = page < total_pages |
| 1038 | has_previous = page > 1 |
| 1039 | |
| 1040 | if total_count == 0: |
| 1041 | return VulnerabilitySearchResponse( |
| 1042 | vulnerabilities=[], |
| 1043 | total_count=0, |
| 1044 | critical_count=0, |
| 1045 | high_count=0, |
| 1046 | medium_count=0, |
| 1047 | low_count=0, |
| 1048 | page=page, |
| 1049 | page_size=page_size, |
| 1050 | total_pages=0, |
| 1051 | has_next=False, |
| 1052 | has_previous=False, |
| 1053 | success=True, |
| 1054 | message="No vulnerabilities found matching the specified criteria", |
| 1055 | filters_applied=filters_applied, |
| 1056 | ) |
| 1057 | |
| 1058 | # Calculate start index |
| 1059 | start_index = (page - 1) * page_size |
| 1060 | |
| 1061 | # Check if pagination exceeds Elasticsearch's 10,000 result window |
| 1062 | if start_index >= 10000: |
| 1063 | logger.warning( |
| 1064 | f"Deep pagination requested (page {page}, start_index {start_index}). " |
| 1065 | f"Elasticsearch limits pagination to 10,000 results. " |
| 1066 | f"Please use more specific filters or export to CSV report.", |
| 1067 | ) |
| 1068 | |
| 1069 | return VulnerabilitySearchResponse( |
| 1070 | vulnerabilities=[], |
| 1071 | total_count=total_count, |
| 1072 | critical_count=severity_counts["Critical"], |
| 1073 | high_count=severity_counts["High"], |
| 1074 | medium_count=severity_counts["Medium"], |
| 1075 | low_count=severity_counts["Low"], |
| 1076 | page=page, |
| 1077 | page_size=page_size, |
| 1078 | total_pages=total_pages, |
| 1079 | has_next=has_next, |
| 1080 | has_previous=has_previous, |
| 1081 | success=False, |
| 1082 | message=( |
| 1083 | f"Deep pagination not supported beyond 10,000 results (requested page {page}, position {start_index}). " |
| 1084 | "Please use more specific filters to narrow down results or use the CSV export feature " |
| 1085 | "for accessing all vulnerabilities. Maximum supported page: 200 (with page_size=50)." |
| 1086 | ), |
| 1087 | filters_applied=filters_applied, |
| 1088 | ) |
| 1089 | |
| 1090 | # Use standard pagination (within Elasticsearch limits) |
| 1091 | search_response = await es_client.search( |
| 1092 | index=",".join(vuln_indices), |
| 1093 | body={ |
| 1094 | "query": es_query, |
| 1095 | "sort": [ |
| 1096 | {"vulnerability.detected_at": {"order": "desc"}}, |
| 1097 | {"vulnerability.severity": {"order": "asc"}}, |
| 1098 | {"_id": {"order": "asc"}}, # Tiebreaker for consistent sorting |
| 1099 | ], |
| 1100 | "from": start_index, |
| 1101 | "size": page_size, |
| 1102 | }, |
| 1103 | ) |
| 1104 | hits = search_response["hits"]["hits"] |
| 1105 | |
| 1106 | # Process the results |
| 1107 | vulnerabilities = [] |
| 1108 | for hit in hits: |
| 1109 | try: |
| 1110 | source = hit["_source"] |
| 1111 | agent_data = source.get("agent", {}) |
| 1112 | agent_hostname = agent_data.get("name", "unknown") |
| 1113 | |
| 1114 | # Get customer code from our mapping |
| 1115 | agent_customer_code = customer_agent_map.get(agent_hostname) |
| 1116 | |
| 1117 | # Process the vulnerability data |
| 1118 | vuln_data = process_wazuh_document(hit) |
| 1119 | |
| 1120 | # Get EPSS score for the CVE (if requested) |
| 1121 | epss_score, epss_percentile = None, None |
| 1122 | if include_epss: |
| 1123 | epss_score, epss_percentile = await get_epss_score_for_cve(vuln_data.cve_id) |
| 1124 | |
| 1125 | vulnerability_item = VulnerabilitySearchItem( |
| 1126 | cve_id=vuln_data.cve_id, |
| 1127 | severity=vuln_data.severity, |
| 1128 | title=vuln_data.title, |
| 1129 | agent_name=agent_hostname, |
| 1130 | customer_code=agent_customer_code, |
| 1131 | references=vuln_data.references, |
| 1132 | detected_at=vuln_data.detected_at, |
| 1133 | published_at=vuln_data.published_at, |
| 1134 | base_score=vuln_data.base_score, |
| 1135 | package_name=vuln_data.package_name, |
| 1136 | package_version=vuln_data.package_version, |
| 1137 | package_architecture=vuln_data.package_architecture, |
| 1138 | epss_score=epss_score, |
| 1139 | epss_percentile=epss_percentile, |
| 1140 | ) |
| 1141 | vulnerabilities.append(vulnerability_item) |
| 1142 | |
| 1143 | except Exception as e: |
| 1144 | logger.error(f"Error processing vulnerability document: {e}") |
| 1145 | continue |
| 1146 | |
| 1147 | # Sort vulnerabilities by EPSS score if included |
| 1148 | if include_epss: |
| 1149 | severity_order = {"Critical": 0, "High": 1, "Medium": 2, "Low": 3} |
| 1150 | |
| 1151 | def get_epss_sort_key(vuln): |
| 1152 | epss_score_val = vuln.epss_score |
| 1153 | if epss_score_val is None: |
| 1154 | epss_float = 0.0 |
| 1155 | else: |
| 1156 | try: |
| 1157 | epss_float = float(epss_score_val) |
| 1158 | except (ValueError, TypeError): |
| 1159 | epss_float = 0.0 |
| 1160 | return ( |
| 1161 | -epss_float, |
| 1162 | severity_order.get(vuln.severity, 4), |
| 1163 | vuln.cve_id, |
| 1164 | ) |
| 1165 | |
| 1166 | vulnerabilities.sort(key=get_epss_sort_key) |
| 1167 | logger.info(f"Sorted {len(vulnerabilities)} vulnerabilities by EPSS score (highest to lowest)") |
| 1168 | |
| 1169 | message = f"Found {len(vulnerabilities)} vulnerabilities on page {page} of {total_pages}" |
| 1170 | if filters_applied: |
| 1171 | message += f" with filters: {filters_applied}" |
| 1172 | if include_epss: |
| 1173 | message += " (sorted by EPSS score, highest to lowest)" |
| 1174 | else: |
| 1175 | message += " (sorted by detection date and severity)" |
| 1176 | |
| 1177 | # Add helpful message when approaching the limit |
| 1178 | if start_index + page_size > 9000: |
| 1179 | message += ( |
| 1180 | ". Note: Approaching pagination limit (10,000 results). Consider using filters or CSV export for complete data access." |
| 1181 | ) |
| 1182 | |
| 1183 | return VulnerabilitySearchResponse( |
| 1184 | vulnerabilities=vulnerabilities, |
| 1185 | total_count=total_count, |
| 1186 | critical_count=severity_counts["Critical"], |
| 1187 | high_count=severity_counts["High"], |
| 1188 | medium_count=severity_counts["Medium"], |
| 1189 | low_count=severity_counts["Low"], |
| 1190 | page=page, |
| 1191 | page_size=page_size, |
| 1192 | total_pages=total_pages, |
| 1193 | has_next=has_next, |
| 1194 | has_previous=has_previous, |
| 1195 | success=True, |
| 1196 | message=message, |
| 1197 | filters_applied=filters_applied, |
| 1198 | ) |
| 1199 | |
| 1200 | except Exception as e: |
| 1201 | logger.error(f"Error searching vulnerabilities from indexer: {e}") |
| 1202 | return VulnerabilitySearchResponse( |
| 1203 | vulnerabilities=[], |
| 1204 | total_count=0, |
| 1205 | critical_count=0, |
| 1206 | high_count=0, |
| 1207 | medium_count=0, |
| 1208 | low_count=0, |
| 1209 | page=page, |
| 1210 | page_size=page_size, |
| 1211 | total_pages=0, |
| 1212 | has_next=False, |
| 1213 | has_previous=False, |
| 1214 | success=False, |
| 1215 | message=f"Failed to search vulnerabilities: {e}", |
| 1216 | filters_applied=filters_applied if "filters_applied" in locals() else {}, |
| 1217 | ) |
| 1218 | finally: |
| 1219 | if es_client: |
| 1220 | try: |
| 1221 | await es_client.close() |
| 1222 | except Exception as close_error: |
| 1223 | logger.warning(f"Error closing Elasticsearch client: {close_error}") |
| 1224 | |
| 1225 | |
| 1226 | async def fetch_all_vulnerabilities_for_export( |
| 1227 | db_session: AsyncSession, |
| 1228 | current_user: User, |
| 1229 | customer_code: str, |
| 1230 | agent_name: Optional[str] = None, |
| 1231 | severity: Optional[str] = None, |
| 1232 | cve_id: Optional[str] = None, |
| 1233 | package_name: Optional[str] = None, |
| 1234 | include_epss: bool = True, |
| 1235 | ) -> List[VulnerabilitySearchItem]: |
| 1236 | """ |
| 1237 | Fetch ALL vulnerabilities for CSV export using scroll API (no pagination limits). |
| 1238 | |
| 1239 | This function is specifically designed for report generation and can handle |
| 1240 | unlimited result sets by using Elasticsearch's scroll API. |
| 1241 | |
| 1242 | Args: |
| 1243 | db_session: Database session for agent lookup |
| 1244 | current_user: Current authenticated user for customer access filtering |
| 1245 | customer_code: Customer code to filter by |
| 1246 | agent_name: Optional agent hostname filter |
| 1247 | severity: Optional severity filter |
| 1248 | cve_id: Optional CVE ID filter |
| 1249 | package_name: Optional package name filter |
| 1250 | include_epss: Whether to include EPSS scores |
| 1251 | |
| 1252 | Returns: |
| 1253 | List of all matching vulnerabilities (no pagination) |
| 1254 | """ |
| 1255 | logger.info( |
| 1256 | f"Fetching ALL vulnerabilities for export with filters: customer_code={customer_code}, " |
| 1257 | f"agent_name={agent_name}, severity={severity}, cve_id={cve_id}, " |
| 1258 | f"package_name={package_name}, include_epss={include_epss}", |
| 1259 | ) |
| 1260 | |
| 1261 | # Apply customer access filtering |
| 1262 | accessible_customers = await customer_access_handler.get_user_accessible_customers(current_user, db_session) |
| 1263 | |
| 1264 | if "*" not in accessible_customers and customer_code not in accessible_customers: |
| 1265 | logger.warning(f"User {current_user.username} denied access to customer {customer_code}") |
| 1266 | return [] |
| 1267 | |
| 1268 | es_client = None |
| 1269 | scroll_id = None |
| 1270 | |
| 1271 | try: |
| 1272 | # Initialize Elasticsearch client |
| 1273 | es_client = await create_wazuh_indexer_client_async("Wazuh-Indexer") |
| 1274 | |
| 1275 | # Get all agents for customer code mapping |
| 1276 | all_agents_query = select(Agents) |
| 1277 | all_agents_result = await db_session.execute(all_agents_query) |
| 1278 | all_agents = all_agents_result.scalars().all() |
| 1279 | |
| 1280 | # Build complete agent hostname to customer code mapping |
| 1281 | customer_agent_map = {} |
| 1282 | for agent in all_agents: |
| 1283 | if agent.hostname: |
| 1284 | customer_agent_map[agent.hostname] = agent.customer_code |
| 1285 | |
| 1286 | # Get agent information for filtering |
| 1287 | agent_hostnames = [] |
| 1288 | query = select(Agents).filter(Agents.customer_code == customer_code) |
| 1289 | |
| 1290 | if agent_name: |
| 1291 | query = query.filter(Agents.hostname == agent_name) |
| 1292 | |
| 1293 | result = await db_session.execute(query) |
| 1294 | agents = result.scalars().all() |
| 1295 | |
| 1296 | if not agents: |
| 1297 | logger.warning(f"No agents found for customer {customer_code}") |
| 1298 | return [] |
| 1299 | |
| 1300 | for agent in agents: |
| 1301 | if agent.hostname: |
| 1302 | agent_hostnames.append(agent.hostname) |
| 1303 | |
| 1304 | # Get vulnerability indices |
| 1305 | vuln_indices = await get_vulnerabilities_indices() |
| 1306 | if not vuln_indices: |
| 1307 | logger.warning("No vulnerability indices found") |
| 1308 | return [] |
| 1309 | |
| 1310 | # Build Elasticsearch query |
| 1311 | es_query = {"bool": {"must": []}} |
| 1312 | |
| 1313 | # Add agent filter |
| 1314 | if agent_hostnames: |
| 1315 | es_query["bool"]["must"].append({"terms": {"agent.name": agent_hostnames}}) |
| 1316 | |
| 1317 | # Add severity filter |
| 1318 | if severity: |
| 1319 | es_query["bool"]["must"].append({"term": {"vulnerability.severity": severity}}) |
| 1320 | |
| 1321 | # Add CVE ID filter |
| 1322 | if cve_id: |
| 1323 | es_query["bool"]["must"].append({"term": {"vulnerability.id": cve_id}}) |
| 1324 | |
| 1325 | # Add package name filter |
| 1326 | if package_name: |
| 1327 | es_query["bool"]["must"].append({"wildcard": {"package.name": f"*{package_name}*"}}) |
| 1328 | |
| 1329 | # Use scroll API for unlimited results |
| 1330 | all_vulnerabilities = [] |
| 1331 | scroll_size = 1000 # Process 1000 at a time |
| 1332 | |
| 1333 | logger.info(f"Starting scroll search across {len(vuln_indices)} indices") |
| 1334 | |
| 1335 | # Initial scroll request |
| 1336 | scroll_response = await es_client.search( |
| 1337 | index=",".join(vuln_indices), |
| 1338 | body={ |
| 1339 | "query": es_query, |
| 1340 | "sort": [ |
| 1341 | {"vulnerability.detected_at": {"order": "desc"}}, |
| 1342 | {"vulnerability.severity": {"order": "asc"}}, |
| 1343 | {"_id": {"order": "asc"}}, |
| 1344 | ], |
| 1345 | }, |
| 1346 | scroll="5m", # Keep scroll context alive for 5 minutes |
| 1347 | size=scroll_size, |
| 1348 | ) |
| 1349 | |
| 1350 | scroll_id = scroll_response["_scroll_id"] |
| 1351 | hits = scroll_response["hits"]["hits"] |
| 1352 | |
| 1353 | logger.info(f"Initial scroll batch: {len(hits)} results") |
| 1354 | |
| 1355 | # Process initial batch |
| 1356 | for hit in hits: |
| 1357 | try: |
| 1358 | source = hit["_source"] |
| 1359 | agent_data = source.get("agent", {}) |
| 1360 | agent_hostname = agent_data.get("name", "unknown") |
| 1361 | agent_customer_code = customer_agent_map.get(agent_hostname) |
| 1362 | |
| 1363 | vuln_data = process_wazuh_document(hit) |
| 1364 | |
| 1365 | # Get EPSS score if requested |
| 1366 | epss_score, epss_percentile = None, None |
| 1367 | if include_epss: |
| 1368 | epss_score, epss_percentile = await get_epss_score_for_cve(vuln_data.cve_id) |
| 1369 | |
| 1370 | vulnerability_item = VulnerabilitySearchItem( |
| 1371 | cve_id=vuln_data.cve_id, |
| 1372 | severity=vuln_data.severity, |
| 1373 | title=vuln_data.title, |
| 1374 | agent_name=agent_hostname, |
| 1375 | customer_code=agent_customer_code, |
| 1376 | references=vuln_data.references, |
| 1377 | detected_at=vuln_data.detected_at, |
| 1378 | published_at=vuln_data.published_at, |
| 1379 | base_score=vuln_data.base_score, |
| 1380 | package_name=vuln_data.package_name, |
| 1381 | package_version=vuln_data.package_version, |
| 1382 | package_architecture=vuln_data.package_architecture, |
| 1383 | epss_score=epss_score, |
| 1384 | epss_percentile=epss_percentile, |
| 1385 | ) |
| 1386 | all_vulnerabilities.append(vulnerability_item) |
| 1387 | |
| 1388 | except Exception as e: |
| 1389 | logger.error(f"Error processing vulnerability document: {e}") |
| 1390 | continue |
| 1391 | |
| 1392 | # Continue scrolling through all results |
| 1393 | scroll_count = 1 |
| 1394 | while len(hits) > 0: |
| 1395 | try: |
| 1396 | scroll_response = await es_client.scroll(scroll_id=scroll_id, scroll="5m") |
| 1397 | scroll_id = scroll_response["_scroll_id"] |
| 1398 | hits = scroll_response["hits"]["hits"] |
| 1399 | |
| 1400 | if not hits: |
| 1401 | break |
| 1402 | |
| 1403 | scroll_count += 1 |
| 1404 | logger.info(f"Scroll batch {scroll_count}: {len(hits)} results (total so far: {len(all_vulnerabilities)})") |
| 1405 | |
| 1406 | # Process batch |
| 1407 | for hit in hits: |
| 1408 | try: |
| 1409 | source = hit["_source"] |
| 1410 | agent_data = source.get("agent", {}) |
| 1411 | agent_hostname = agent_data.get("name", "unknown") |
| 1412 | agent_customer_code = customer_agent_map.get(agent_hostname) |
| 1413 | |
| 1414 | vuln_data = process_wazuh_document(hit) |
| 1415 | |
| 1416 | # Get EPSS score if requested |
| 1417 | epss_score, epss_percentile = None, None |
| 1418 | if include_epss: |
| 1419 | epss_score, epss_percentile = await get_epss_score_for_cve(vuln_data.cve_id) |
| 1420 | |
| 1421 | vulnerability_item = VulnerabilitySearchItem( |
| 1422 | cve_id=vuln_data.cve_id, |
| 1423 | severity=vuln_data.severity, |
| 1424 | title=vuln_data.title, |
| 1425 | agent_name=agent_hostname, |
| 1426 | customer_code=agent_customer_code, |
| 1427 | references=vuln_data.references, |
| 1428 | detected_at=vuln_data.detected_at, |
| 1429 | published_at=vuln_data.published_at, |
| 1430 | base_score=vuln_data.base_score, |
| 1431 | package_name=vuln_data.package_name, |
| 1432 | package_version=vuln_data.package_version, |
| 1433 | package_architecture=vuln_data.package_architecture, |
| 1434 | epss_score=epss_score, |
| 1435 | epss_percentile=epss_percentile, |
| 1436 | ) |
| 1437 | all_vulnerabilities.append(vulnerability_item) |
| 1438 | |
| 1439 | except Exception as e: |
| 1440 | logger.error(f"Error processing vulnerability document: {e}") |
| 1441 | continue |
| 1442 | |
| 1443 | except Exception as scroll_error: |
| 1444 | logger.error(f"Error during scroll: {scroll_error}") |
| 1445 | break |
| 1446 | |
| 1447 | logger.info(f"Successfully fetched {len(all_vulnerabilities)} total vulnerabilities using scroll API") |
| 1448 | |
| 1449 | return all_vulnerabilities |
| 1450 | |
| 1451 | except Exception as e: |
| 1452 | logger.error(f"Error fetching vulnerabilities for export: {e}") |
| 1453 | raise |
| 1454 | |
| 1455 | finally: |
| 1456 | # Always clear the scroll context |
| 1457 | if scroll_id and es_client: |
| 1458 | try: |
| 1459 | await es_client.clear_scroll(scroll_id=scroll_id) |
| 1460 | logger.info("Cleared scroll context") |
| 1461 | except Exception as clear_error: |
| 1462 | logger.warning(f"Could not clear scroll context: {clear_error}") |
| 1463 | |
| 1464 | # Close ES client |
| 1465 | if es_client: |
| 1466 | try: |
| 1467 | await es_client.close() |
| 1468 | except Exception as close_error: |
| 1469 | logger.warning(f"Error closing Elasticsearch client: {close_error}") |
| 1470 | |
| 1471 | |
| 1472 | async def generate_vulnerability_csv_report( |
| 1473 | db_session: AsyncSession, |
| 1474 | current_user: User, |
| 1475 | request: VulnerabilityReportGenerateRequest, |
| 1476 | report_id: Optional[int] = None, |
| 1477 | ) -> VulnerabilityReportGenerateResponse: |
| 1478 | """ |
| 1479 | Generate a CSV vulnerability report for a specific customer and store it in MinIO |
| 1480 | |
| 1481 | Args: |
| 1482 | db_session: Database session |
| 1483 | current_user: Current authenticated user |
| 1484 | request: Report generation request with filters |
| 1485 | report_id: Optional existing report ID (for background task updates) |
| 1486 | |
| 1487 | Returns: |
| 1488 | VulnerabilityReportGenerateResponse with report details |
| 1489 | """ |
| 1490 | try: |
| 1491 | # Verify customer access |
| 1492 | accessible_customers = await customer_access_handler.get_user_accessible_customers(current_user, db_session) |
| 1493 | |
| 1494 | if "*" not in accessible_customers and request.customer_code not in accessible_customers: |
| 1495 | # If we have a report_id, update it to failed status |
| 1496 | if report_id: |
| 1497 | stmt = select(VulnerabilityReport).filter(VulnerabilityReport.id == report_id) |
| 1498 | result = await db_session.execute(stmt) |
| 1499 | report = result.scalars().first() |
| 1500 | if report: |
| 1501 | report.status = "failed" |
| 1502 | report.error_message = "Insufficient permissions" |
| 1503 | await db_session.commit() |
| 1504 | |
| 1505 | return VulnerabilityReportGenerateResponse( |
| 1506 | success=False, |
| 1507 | message=f"Access denied to customer {request.customer_code}", |
| 1508 | error="Insufficient permissions", |
| 1509 | ) |
| 1510 | |
| 1511 | # Verify customer exists |
| 1512 | customer_result = await db_session.execute(select(Customers).filter(Customers.customer_code == request.customer_code)) |
| 1513 | customer = customer_result.scalars().first() |
| 1514 | |
| 1515 | if not customer: |
| 1516 | # If we have a report_id, update it to failed status |
| 1517 | if report_id: |
| 1518 | stmt = select(VulnerabilityReport).filter(VulnerabilityReport.id == report_id) |
| 1519 | result = await db_session.execute(stmt) |
| 1520 | report = result.scalars().first() |
| 1521 | if report: |
| 1522 | report.status = "failed" |
| 1523 | report.error_message = "Customer not found" |
| 1524 | await db_session.commit() |
| 1525 | |
| 1526 | return VulnerabilityReportGenerateResponse( |
| 1527 | success=False, |
| 1528 | message=f"Customer {request.customer_code} not found", |
| 1529 | error="Customer not found", |
| 1530 | ) |
| 1531 | |
| 1532 | # If report_id exists, get the existing report name and object_key |
| 1533 | # Otherwise generate new ones |
| 1534 | if report_id: |
| 1535 | stmt = select(VulnerabilityReport).filter(VulnerabilityReport.id == report_id) |
| 1536 | result = await db_session.execute(stmt) |
| 1537 | existing_report = result.scalars().first() |
| 1538 | |
| 1539 | if not existing_report: |
| 1540 | return VulnerabilityReportGenerateResponse( |
| 1541 | success=False, |
| 1542 | message=f"Report ID {report_id} not found", |
| 1543 | error="Report not found", |
| 1544 | ) |
| 1545 | |
| 1546 | # Use existing report name and paths |
| 1547 | report_name = existing_report.report_name |
| 1548 | file_name = existing_report.file_name |
| 1549 | object_key = existing_report.object_key |
| 1550 | bucket_name = existing_report.bucket_name |
| 1551 | |
| 1552 | logger.info(f"Generating vulnerability report for existing record: {report_name} (ID: {report_id})") |
| 1553 | else: |
| 1554 | # Generate new report name for synchronous generation |
| 1555 | timestamp = datetime.utcnow().strftime("%Y%m%d_%H%M%S") |
| 1556 | report_name = request.report_name or f"vulnerability_report_{timestamp}" |
| 1557 | file_name = f"{report_name}.csv" |
| 1558 | object_key = f"{request.customer_code}/{file_name}" |
| 1559 | bucket_name = "vulnerability-reports" |
| 1560 | |
| 1561 | logger.info(f"Generating new vulnerability report: {report_name}") |
| 1562 | |
| 1563 | # Fetch ALL vulnerabilities using scroll API (no pagination limits) |
| 1564 | all_vulnerabilities = await fetch_all_vulnerabilities_for_export( |
| 1565 | db_session=db_session, |
| 1566 | current_user=current_user, |
| 1567 | customer_code=request.customer_code, |
| 1568 | agent_name=request.agent_name, |
| 1569 | severity=request.severity, |
| 1570 | cve_id=request.cve_id, |
| 1571 | package_name=request.package_name, |
| 1572 | include_epss=request.include_epss, |
| 1573 | ) |
| 1574 | |
| 1575 | logger.info(f"Fetched {len(all_vulnerabilities)} vulnerabilities for report") |
| 1576 | |
| 1577 | if not all_vulnerabilities: |
| 1578 | # If we have a report_id, update it to failed status |
| 1579 | if report_id: |
| 1580 | stmt = select(VulnerabilityReport).filter(VulnerabilityReport.id == report_id) |
| 1581 | result = await db_session.execute(stmt) |
| 1582 | report = result.scalars().first() |
| 1583 | if report: |
| 1584 | report.status = "failed" |
| 1585 | report.error_message = "No vulnerabilities found matching criteria" |
| 1586 | await db_session.commit() |
| 1587 | |
| 1588 | return VulnerabilityReportGenerateResponse( |
| 1589 | success=False, |
| 1590 | message="No vulnerabilities found matching the specified criteria", |
| 1591 | error="No data to export", |
| 1592 | ) |
| 1593 | |
| 1594 | # Generate CSV content |
| 1595 | csv_buffer = io.StringIO() |
| 1596 | csv_writer = csv.writer(csv_buffer) |
| 1597 | |
| 1598 | # Write headers |
| 1599 | headers = [ |
| 1600 | "CVE ID", |
| 1601 | "Severity", |
| 1602 | "Title", |
| 1603 | "Agent Name", |
| 1604 | "Customer Code", |
| 1605 | "Package Name", |
| 1606 | "Package Version", |
| 1607 | "Package Architecture", |
| 1608 | "Detected At", |
| 1609 | "Published At", |
| 1610 | "Base Score", |
| 1611 | ] |
| 1612 | |
| 1613 | if request.include_epss: |
| 1614 | headers.extend(["EPSS Score", "EPSS Percentile"]) |
| 1615 | |
| 1616 | headers.append("References") |
| 1617 | csv_writer.writerow(headers) |
| 1618 | |
| 1619 | # Write data rows |
| 1620 | for vuln in all_vulnerabilities: |
| 1621 | row = [ |
| 1622 | vuln.cve_id, |
| 1623 | vuln.severity, |
| 1624 | vuln.title, |
| 1625 | vuln.agent_name, |
| 1626 | vuln.customer_code or "", |
| 1627 | vuln.package_name or "", |
| 1628 | vuln.package_version or "", |
| 1629 | vuln.package_architecture or "", |
| 1630 | vuln.detected_at.isoformat() if vuln.detected_at else "", |
| 1631 | vuln.published_at.isoformat() if vuln.published_at else "", |
| 1632 | vuln.base_score or "", |
| 1633 | ] |
| 1634 | |
| 1635 | if request.include_epss: |
| 1636 | row.extend( |
| 1637 | [ |
| 1638 | vuln.epss_score or "", |
| 1639 | vuln.epss_percentile or "", |
| 1640 | ], |
| 1641 | ) |
| 1642 | |
| 1643 | row.append(vuln.references or "") |
| 1644 | csv_writer.writerow(row) |
| 1645 | |
| 1646 | # Get CSV content as bytes |
| 1647 | csv_content = csv_buffer.getvalue().encode("utf-8") |
| 1648 | csv_buffer.close() |
| 1649 | |
| 1650 | # Calculate file hash |
| 1651 | file_hash = hashlib.sha256(csv_content).hexdigest() |
| 1652 | |
| 1653 | # Store in MinIO using the paths we determined earlier |
| 1654 | minio_result = await store_file_in_minio( |
| 1655 | file_content=csv_content, |
| 1656 | bucket_name=bucket_name, |
| 1657 | object_key=object_key, |
| 1658 | content_type="text/csv", |
| 1659 | ) |
| 1660 | |
| 1661 | if not minio_result["success"]: |
| 1662 | # If we have a report_id, update it to failed status |
| 1663 | if report_id: |
| 1664 | stmt = select(VulnerabilityReport).filter(VulnerabilityReport.id == report_id) |
| 1665 | result = await db_session.execute(stmt) |
| 1666 | report = result.scalars().first() |
| 1667 | if report: |
| 1668 | report.status = "failed" |
| 1669 | report.error_message = minio_result.get("error", "Unknown error") |
| 1670 | await db_session.commit() |
| 1671 | |
| 1672 | return VulnerabilityReportGenerateResponse( |
| 1673 | success=False, |
| 1674 | message="Failed to store report in MinIO", |
| 1675 | error=minio_result.get("error", "Unknown error"), |
| 1676 | ) |
| 1677 | |
| 1678 | # Build filters JSON |
| 1679 | filters = {} |
| 1680 | if request.agent_name: |
| 1681 | filters["agent_name"] = request.agent_name |
| 1682 | if request.severity: |
| 1683 | filters["severity"] = request.severity |
| 1684 | if request.cve_id: |
| 1685 | filters["cve_id"] = request.cve_id |
| 1686 | if request.package_name: |
| 1687 | filters["package_name"] = request.package_name |
| 1688 | filters["include_epss"] = request.include_epss |
| 1689 | |
| 1690 | # Check if we're updating an existing report or creating a new one |
| 1691 | if report_id: |
| 1692 | # Update existing report (background task scenario) |
| 1693 | stmt = select(VulnerabilityReport).filter(VulnerabilityReport.id == report_id) |
| 1694 | result = await db_session.execute(stmt) |
| 1695 | report_record = result.scalars().first() |
| 1696 | |
| 1697 | if report_record: |
| 1698 | report_record.file_size = len(csv_content) |
| 1699 | report_record.file_hash = file_hash |
| 1700 | report_record.total_vulnerabilities = len(all_vulnerabilities) |
| 1701 | report_record.critical_count = sum(1 for v in all_vulnerabilities if v.severity == "Critical") |
| 1702 | report_record.high_count = sum(1 for v in all_vulnerabilities if v.severity == "High") |
| 1703 | report_record.medium_count = sum(1 for v in all_vulnerabilities if v.severity == "Medium") |
| 1704 | report_record.low_count = sum(1 for v in all_vulnerabilities if v.severity == "Low") |
| 1705 | report_record.status = "completed" |
| 1706 | report_record.error_message = None |
| 1707 | |
| 1708 | await db_session.commit() |
| 1709 | await db_session.refresh(report_record) |
| 1710 | |
| 1711 | logger.info(f"Successfully updated vulnerability report: {report_name} (ID: {report_id})") |
| 1712 | else: |
| 1713 | logger.error(f"Report ID {report_id} not found for update") |
| 1714 | return VulnerabilityReportGenerateResponse( |
| 1715 | success=False, |
| 1716 | message=f"Report ID {report_id} not found", |
| 1717 | error="Report not found", |
| 1718 | ) |
| 1719 | else: |
| 1720 | # Create new database record (synchronous scenario) |
| 1721 | report_record = VulnerabilityReport( |
| 1722 | report_name=report_name, |
| 1723 | customer_code=request.customer_code, |
| 1724 | bucket_name=bucket_name, |
| 1725 | object_key=object_key, |
| 1726 | file_name=file_name, |
| 1727 | file_size=len(csv_content), |
| 1728 | file_hash=file_hash, |
| 1729 | generated_by=current_user.id, |
| 1730 | filters_json=json.dumps(filters), |
| 1731 | total_vulnerabilities=len(all_vulnerabilities), |
| 1732 | critical_count=sum(1 for v in all_vulnerabilities if v.severity == "Critical"), |
| 1733 | high_count=sum(1 for v in all_vulnerabilities if v.severity == "High"), |
| 1734 | medium_count=sum(1 for v in all_vulnerabilities if v.severity == "Medium"), |
| 1735 | low_count=sum(1 for v in all_vulnerabilities if v.severity == "Low"), |
| 1736 | status="completed", |
| 1737 | ) |
| 1738 | |
| 1739 | db_session.add(report_record) |
| 1740 | await db_session.commit() |
| 1741 | await db_session.refresh(report_record) |
| 1742 | |
| 1743 | logger.info(f"Successfully generated vulnerability report: {report_name}") |
| 1744 | |
| 1745 | # Build response |
| 1746 | report_response = VulnerabilityReportResponse( |
| 1747 | id=report_record.id, |
| 1748 | report_name=report_record.report_name, |
| 1749 | customer_code=report_record.customer_code, |
| 1750 | file_name=report_record.file_name, |
| 1751 | file_size=report_record.file_size, |
| 1752 | generated_at=report_record.generated_at, |
| 1753 | generated_by=report_record.generated_by, |
| 1754 | total_vulnerabilities=report_record.total_vulnerabilities, |
| 1755 | critical_count=report_record.critical_count, |
| 1756 | high_count=report_record.high_count, |
| 1757 | medium_count=report_record.medium_count, |
| 1758 | low_count=report_record.low_count, |
| 1759 | filters_applied=json.loads(report_record.filters_json or "{}"), |
| 1760 | status=report_record.status, |
| 1761 | download_url=f"/api/v1/vulnerabilities/reports/{report_record.id}/download", |
| 1762 | ) |
| 1763 | |
| 1764 | return VulnerabilityReportGenerateResponse( |
| 1765 | success=True, |
| 1766 | message=f"Successfully generated report with {len(all_vulnerabilities)} vulnerabilities", |
| 1767 | report=report_response, |
| 1768 | ) |
| 1769 | |
| 1770 | except Exception as e: |
| 1771 | logger.error(f"Error generating vulnerability report: {e}") |
| 1772 | |
| 1773 | # If we have a report_id, update it to failed status |
| 1774 | if report_id: |
| 1775 | try: |
| 1776 | stmt = select(VulnerabilityReport).filter(VulnerabilityReport.id == report_id) |
| 1777 | result = await db_session.execute(stmt) |
| 1778 | report = result.scalars().first() |
| 1779 | if report: |
| 1780 | report.status = "failed" |
| 1781 | report.error_message = str(e) |
| 1782 | await db_session.commit() |
| 1783 | except Exception as update_error: |
| 1784 | logger.error(f"Failed to update report status: {update_error}") |
| 1785 | |
| 1786 | return VulnerabilityReportGenerateResponse( |
| 1787 | success=False, |
| 1788 | message="Failed to generate vulnerability report", |
| 1789 | error=str(e), |
| 1790 | ) |
| 1791 | |
| 1792 | |
| 1793 | async def list_vulnerability_reports( |
| 1794 | db_session: AsyncSession, |
| 1795 | current_user: User, |
| 1796 | customer_code: Optional[str] = None, |
| 1797 | ) -> VulnerabilityReportListResponse: |
| 1798 | """ |
| 1799 | List available vulnerability reports |
| 1800 | |
| 1801 | Args: |
| 1802 | db_session: Database session |
| 1803 | current_user: Current authenticated user |
| 1804 | customer_code: Optional filter by customer code |
| 1805 | |
| 1806 | Returns: |
| 1807 | VulnerabilityReportListResponse with list of reports |
| 1808 | """ |
| 1809 | try: |
| 1810 | # Get accessible customers |
| 1811 | accessible_customers = await customer_access_handler.get_user_accessible_customers(current_user, db_session) |
| 1812 | |
| 1813 | # Build query |
| 1814 | query = select(VulnerabilityReport).order_by(desc(VulnerabilityReport.generated_at)) |
| 1815 | |
| 1816 | # Apply customer filtering |
| 1817 | if "*" not in accessible_customers: |
| 1818 | query = query.filter(VulnerabilityReport.customer_code.in_(accessible_customers)) |
| 1819 | |
| 1820 | if customer_code: |
| 1821 | if "*" not in accessible_customers and customer_code not in accessible_customers: |
| 1822 | return VulnerabilityReportListResponse( |
| 1823 | reports=[], |
| 1824 | total_count=0, |
| 1825 | success=True, |
| 1826 | message=f"Access denied to customer {customer_code}", |
| 1827 | ) |
| 1828 | query = query.filter(VulnerabilityReport.customer_code == customer_code) |
| 1829 | |
| 1830 | result = await db_session.execute(query) |
| 1831 | reports = result.scalars().all() |
| 1832 | |
| 1833 | report_list = [] |
| 1834 | for report in reports: |
| 1835 | report_response = VulnerabilityReportResponse( |
| 1836 | id=report.id, |
| 1837 | report_name=report.report_name, |
| 1838 | customer_code=report.customer_code, |
| 1839 | file_name=report.file_name, |
| 1840 | file_size=report.file_size, |
| 1841 | generated_at=report.generated_at, |
| 1842 | generated_by=report.generated_by, |
| 1843 | total_vulnerabilities=report.total_vulnerabilities, |
| 1844 | critical_count=report.critical_count, |
| 1845 | high_count=report.high_count, |
| 1846 | medium_count=report.medium_count, |
| 1847 | low_count=report.low_count, |
| 1848 | filters_applied=json.loads(report.filters_json or "{}"), |
| 1849 | status=report.status, |
| 1850 | download_url=f"/api/v1/vulnerabilities/reports/{report.id}/download", |
| 1851 | ) |
| 1852 | report_list.append(report_response) |
| 1853 | |
| 1854 | return VulnerabilityReportListResponse( |
| 1855 | reports=report_list, |
| 1856 | total_count=len(report_list), |
| 1857 | success=True, |
| 1858 | message=f"Found {len(report_list)} vulnerability reports", |
| 1859 | ) |
| 1860 | |
| 1861 | except Exception as e: |
| 1862 | logger.error(f"Error listing vulnerability reports: {e}") |
| 1863 | return VulnerabilityReportListResponse( |
| 1864 | reports=[], |
| 1865 | total_count=0, |
| 1866 | success=False, |
| 1867 | message=f"Failed to list reports: {e}", |
| 1868 | ) |
| 1869 | |
| 1870 | |
| 1871 | async def get_vulnerability_report_download( |
| 1872 | db_session: AsyncSession, |
| 1873 | current_user: User, |
| 1874 | report_id: int, |
| 1875 | ) -> Dict[str, Any]: |
| 1876 | """ |
| 1877 | Get vulnerability report for download |
| 1878 | |
| 1879 | Args: |
| 1880 | db_session: Database session |
| 1881 | current_user: Current authenticated user |
| 1882 | report_id: Report ID to download |
| 1883 | |
| 1884 | Returns: |
| 1885 | Dict with file_content, file_name, and content_type |
| 1886 | """ |
| 1887 | try: |
| 1888 | # Get report record |
| 1889 | result = await db_session.execute(select(VulnerabilityReport).filter(VulnerabilityReport.id == report_id)) |
| 1890 | report = result.scalars().first() |
| 1891 | |
| 1892 | if not report: |
| 1893 | raise HTTPException(status_code=404, detail="Report not found") |
| 1894 | |
| 1895 | # Verify customer access |
| 1896 | accessible_customers = await customer_access_handler.get_user_accessible_customers(current_user, db_session) |
| 1897 | |
| 1898 | if "*" not in accessible_customers and report.customer_code not in accessible_customers: |
| 1899 | raise HTTPException(status_code=403, detail="Access denied to this report") |
| 1900 | |
| 1901 | # Retrieve file from MinIO |
| 1902 | from app.data_store.data_store_operations import retrieve_file_from_minio |
| 1903 | |
| 1904 | file_data = await retrieve_file_from_minio( |
| 1905 | bucket_name=report.bucket_name, |
| 1906 | object_key=report.object_key, |
| 1907 | ) |
| 1908 | |
| 1909 | if not file_data["success"]: |
| 1910 | raise HTTPException(status_code=500, detail="Failed to retrieve report file") |
| 1911 | |
| 1912 | return { |
| 1913 | "file_content": file_data["file_content"], |
| 1914 | "file_name": report.file_name, |
| 1915 | "content_type": "text/csv", |
| 1916 | } |
| 1917 | |
| 1918 | except HTTPException: |
| 1919 | raise |
| 1920 | except Exception as e: |
| 1921 | logger.error(f"Error retrieving vulnerability report: {e}") |
| 1922 | raise HTTPException(status_code=500, detail=f"Failed to retrieve report: {e}") |
| 1923 | |
| 1924 | |
| 1925 | async def delete_vulnerability_report( |
| 1926 | db_session: AsyncSession, |
| 1927 | current_user: User, |
| 1928 | report_id: int, |
| 1929 | ) -> Dict[str, Any]: |
| 1930 | """ |
| 1931 | Delete a vulnerability report and its associated file from MinIO. |
| 1932 | |
| 1933 | Args: |
| 1934 | db_session: Database session |
| 1935 | current_user: Current authenticated user |
| 1936 | report_id: ID of the report to delete |
| 1937 | |
| 1938 | Returns: |
| 1939 | Dict with success status and details |
| 1940 | """ |
| 1941 | from app.data_store.data_store_operations import delete_file_from_minio |
| 1942 | |
| 1943 | try: |
| 1944 | # Get the report record |
| 1945 | stmt = select(VulnerabilityReport).filter(VulnerabilityReport.id == report_id) |
| 1946 | result = await db_session.execute(stmt) |
| 1947 | report = result.scalars().first() |
| 1948 | |
| 1949 | if not report: |
| 1950 | return { |
| 1951 | "success": False, |
| 1952 | "error": f"Report with ID {report_id} not found", |
| 1953 | } |
| 1954 | |
| 1955 | # Verify customer access |
| 1956 | accessible_customers = await customer_access_handler.get_user_accessible_customers(current_user, db_session) |
| 1957 | |
| 1958 | if "*" not in accessible_customers and report.customer_code not in accessible_customers: |
| 1959 | return { |
| 1960 | "success": False, |
| 1961 | "error": f"Access denied to delete report for customer {report.customer_code}", |
| 1962 | } |
| 1963 | |
| 1964 | logger.info(f"Deleting vulnerability report ID {report_id} for customer {report.customer_code}") |
| 1965 | |
| 1966 | # Delete file from MinIO |
| 1967 | minio_result = await delete_file_from_minio( |
| 1968 | bucket_name=report.bucket_name, |
| 1969 | object_key=report.object_key, |
| 1970 | ) |
| 1971 | |
| 1972 | if not minio_result["success"]: |
| 1973 | logger.warning( |
| 1974 | f"Failed to delete file from MinIO for report {report_id}: {minio_result.get('error')}. " |
| 1975 | "Proceeding with database deletion.", |
| 1976 | ) |
| 1977 | |
| 1978 | # Store report details before deletion |
| 1979 | report_name = report.report_name |
| 1980 | customer_code = report.customer_code |
| 1981 | |
| 1982 | # Delete database record |
| 1983 | await db_session.delete(report) |
| 1984 | await db_session.commit() |
| 1985 | |
| 1986 | logger.info(f"Successfully deleted vulnerability report ID {report_id}") |
| 1987 | |
| 1988 | return { |
| 1989 | "success": True, |
| 1990 | "message": f"Report '{report_name}' deleted successfully", |
| 1991 | "report_id": report_id, |
| 1992 | "report_name": report_name, |
| 1993 | "customer_code": customer_code, |
| 1994 | } |
| 1995 | |
| 1996 | except Exception as e: |
| 1997 | logger.error(f"Error deleting vulnerability report {report_id}: {e}") |
| 1998 | return { |
| 1999 | "success": False, |
| 2000 | "error": str(e), |
| 2001 | } |