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
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
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,
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
1215
message=f"Unexpected error occurred: {e}",
1216
filters_applied=filters_applied if "filters_applied" in locals() else {},
1217
)
1218
+
1219
+
1220
+async def generate_vulnerability_csv_report(
1221
+ db_session: AsyncSession,
1222
+ current_user: User,
1223
+ request: VulnerabilityReportGenerateRequest,
1224
+ report_id: Optional[int] = None,
1225
+) -> VulnerabilityReportGenerateResponse:
1226
+ """
1227
+ Generate a CSV vulnerability report for a specific customer and store it in MinIO
1228
+
1229
+ Args:
1230
+ db_session: Database session
1231
+ current_user: Current authenticated user
1232
+ request: Report generation request with filters
1233
+ report_id: Optional existing report ID (for background task updates)
1234
+
1235
+ Returns:
1236
+ VulnerabilityReportGenerateResponse with report details
1237
+ """
1238
+ try:
1239
+ # Verify customer access
1240
+ accessible_customers = await customer_access_handler.get_user_accessible_customers(current_user, db_session)
1241
+
1242
+ if "*" not in accessible_customers and request.customer_code not in accessible_customers:
1243
+ # If we have a report_id, update it to failed status
1244
+ if report_id:
1245
+ stmt = select(VulnerabilityReport).filter(VulnerabilityReport.id == report_id)
1246
+ result = await db_session.execute(stmt)
1247
+ report = result.scalars().first()
1248
+ if report:
1249
+ report.status = "failed"
1250
+ report.error_message = "Insufficient permissions"
1251
+ await db_session.commit()
1252
+
1253
+ return VulnerabilityReportGenerateResponse(
1254
+ success=False,
1255
+ message=f"Access denied to customer {request.customer_code}",
1256
+ error="Insufficient permissions",
1257
+ )
1258
+
1259
+ # Verify customer exists
1260
+ customer_result = await db_session.execute(select(Customers).filter(Customers.customer_code == request.customer_code))
1261
+ customer = customer_result.scalars().first()
1262
+
1263
+ if not customer:
1264
+ # If we have a report_id, update it to failed status
1265
+ if report_id:
1266
+ stmt = select(VulnerabilityReport).filter(VulnerabilityReport.id == report_id)
1267
+ result = await db_session.execute(stmt)
1268
+ report = result.scalars().first()
1269
+ if report:
1270
+ report.status = "failed"
1271
+ report.error_message = "Customer not found"
1272
+ await db_session.commit()
1273
+
1274
+ return VulnerabilityReportGenerateResponse(
1275
+ success=False,
1276
+ message=f"Customer {request.customer_code} not found",
1277
+ error="Customer not found",
1278
+ )
1279
+
1280
+ logger.info(f"Generating vulnerability report for customer: {request.customer_code}")
1281
+
1282
+ # Fetch ALL vulnerabilities (no pagination)
1283
+ all_vulnerabilities = []
1284
+ page = 1
1285
+ page_size = 1000 # Large page size for efficiency
1286
+
1287
+ while True:
1288
+ search_result = await search_vulnerabilities_from_indexer(
1289
+ db_session=db_session,
1290
+ current_user=current_user,
1291
+ customer_code=request.customer_code,
1292
+ agent_name=request.agent_name,
1293
+ severity=request.severity,
1294
+ cve_id=request.cve_id,
1295
+ package_name=request.package_name,
1296
+ page=page,
1297
+ page_size=page_size,
1298
+ include_epss=request.include_epss,
1299
+ )
1300
+
1301
+ if not search_result.success:
1302
+ # If we have a report_id, update it to failed status
1303
+ if report_id:
1304
+ stmt = select(VulnerabilityReport).filter(VulnerabilityReport.id == report_id)
1305
+ result = await db_session.execute(stmt)
1306
+ report = result.scalars().first()
1307
+ if report:
1308
+ report.status = "failed"
1309
+ report.error_message = search_result.message
1310
+ await db_session.commit()
1311
+
1312
+ return VulnerabilityReportGenerateResponse(
1313
+ success=False,
1314
+ message="Failed to fetch vulnerability data",
1315
+ error=search_result.message,
1316
+ )
1317
+
1318
+ all_vulnerabilities.extend(search_result.vulnerabilities)
1319
+
1320
+ if not search_result.has_next:
1321
+ break
1322
+
1323
+ page += 1
1324
+
1325
+ logger.info(f"Fetched {len(all_vulnerabilities)} vulnerabilities for report")
1326
+
1327
+ # Generate CSV content
1328
+ csv_buffer = io.StringIO()
1329
+ csv_writer = csv.writer(csv_buffer)
1330
+
1331
+ # Write headers
1332
+ headers = [
1333
+ "CVE ID",
1334
+ "Severity",
1335
+ "Title",
1336
+ "Agent Name",
1337
+ "Customer Code",
1338
+ "Package Name",
1339
+ "Package Version",
1340
+ "Package Architecture",
1341
+ "Detected At",
1342
+ "Published At",
1343
+ "Base Score",
1344
+ ]
1345
+
1346
+ if request.include_epss:
1347
+ headers.extend(["EPSS Score", "EPSS Percentile"])
1348
+
1349
+ headers.append("References")
1350
+ csv_writer.writerow(headers)
1351
+
1352
+ # Write data rows
1353
+ for vuln in all_vulnerabilities:
1354
+ row = [
1355
+ vuln.cve_id,
1356
+ vuln.severity,
1357
+ vuln.title,
1358
+ vuln.agent_name,
1359
+ vuln.customer_code or "",
1360
+ vuln.package_name or "",
1361
+ vuln.package_version or "",
1362
+ vuln.package_architecture or "",
1363
+ vuln.detected_at.isoformat() if vuln.detected_at else "",
1364
+ vuln.published_at.isoformat() if vuln.published_at else "",
1365
+ vuln.base_score or "",
1366
+ ]
1367
+
1368
+ if request.include_epss:
1369
+ row.extend(
1370
+ [
1371
+ vuln.epss_score or "",
1372
+ vuln.epss_percentile or "",
1373
+ ],
1374
+ )
1375
+
1376
+ row.append(vuln.references or "")
1377
+ csv_writer.writerow(row)
1378
+
1379
+ # Get CSV content as bytes
1380
+ csv_content = csv_buffer.getvalue().encode("utf-8")
1381
+ csv_buffer.close()
1382
+
1383
+ # Generate file name
1384
+ timestamp = datetime.utcnow().strftime("%Y%m%d_%H%M%S")
1385
+ report_name = request.report_name or f"vulnerability_report_{timestamp}"
1386
+ file_name = f"{report_name}.csv"
1387
+
1388
+ # Calculate file hash
1389
+ file_hash = hashlib.sha256(csv_content).hexdigest()
1390
+
1391
+ # Store in MinIO
1392
+ object_key = f"{request.customer_code}/{file_name}"
1393
+ bucket_name = "vulnerability-reports"
1394
+
1395
+ minio_result = await store_file_in_minio(
1396
+ file_content=csv_content,
1397
+ bucket_name=bucket_name,
1398
+ object_key=object_key,
1399
+ content_type="text/csv",
1400
+ )
1401
+
1402
+ if not minio_result["success"]:
1403
+ # If we have a report_id, update it to failed status
1404
+ if report_id:
1405
+ stmt = select(VulnerabilityReport).filter(VulnerabilityReport.id == report_id)
1406
+ result = await db_session.execute(stmt)
1407
+ report = result.scalars().first()
1408
+ if report:
1409
+ report.status = "failed"
1410
+ report.error_message = minio_result.get("error", "Unknown error")
1411
+ await db_session.commit()
1412
+
1413
+ return VulnerabilityReportGenerateResponse(
1414
+ success=False,
1415
+ message="Failed to store report in MinIO",
1416
+ error=minio_result.get("error", "Unknown error"),
1417
+ )
1418
+
1419
+ # Build filters JSON
1420
+ filters = {}
1421
+ if request.agent_name:
1422
+ filters["agent_name"] = request.agent_name
1423
+ if request.severity:
1424
+ filters["severity"] = request.severity
1425
+ if request.cve_id:
1426
+ filters["cve_id"] = request.cve_id
1427
+ if request.package_name:
1428
+ filters["package_name"] = request.package_name
1429
+ filters["include_epss"] = request.include_epss
1430
+
1431
+ # Check if we're updating an existing report or creating a new one
1432
+ if report_id:
1433
+ # Update existing report (background task scenario)
1434
+ stmt = select(VulnerabilityReport).filter(VulnerabilityReport.id == report_id)
1435
+ result = await db_session.execute(stmt)
1436
+ report_record = result.scalars().first()
1437
+
1438
+ if report_record:
1439
+ report_record.file_size = len(csv_content)
1440
+ report_record.file_hash = file_hash
1441
+ report_record.total_vulnerabilities = len(all_vulnerabilities)
1442
+ report_record.critical_count = sum(1 for v in all_vulnerabilities if v.severity == "Critical")
1443
+ report_record.high_count = sum(1 for v in all_vulnerabilities if v.severity == "High")
1444
+ report_record.medium_count = sum(1 for v in all_vulnerabilities if v.severity == "Medium")
1445
+ report_record.low_count = sum(1 for v in all_vulnerabilities if v.severity == "Low")
1446
+ report_record.status = "completed"
1447
+ report_record.error_message = None
1448
+
1449
+ await db_session.commit()
1450
+ await db_session.refresh(report_record)
1451
+
1452
+ logger.info(f"Successfully updated vulnerability report: {report_name} (ID: {report_id})")
1453
+ else:
1454
+ logger.error(f"Report ID {report_id} not found for update")
1455
+ return VulnerabilityReportGenerateResponse(
1456
+ success=False,
1457
+ message=f"Report ID {report_id} not found",
1458
+ error="Report not found",
1459
+ )
1460
+ else:
1461
+ # Create new database record (synchronous scenario)
1462
+ report_record = VulnerabilityReport(
1463
+ report_name=report_name,
1464
+ customer_code=request.customer_code,
1465
+ bucket_name=bucket_name,
1466
+ object_key=object_key,
1467
+ file_name=file_name,
1468
+ file_size=len(csv_content),
1469
+ file_hash=file_hash,
1470
+ generated_by=current_user.id,
1471
+ filters_json=json.dumps(filters),
1472
+ total_vulnerabilities=len(all_vulnerabilities),
1473
+ critical_count=sum(1 for v in all_vulnerabilities if v.severity == "Critical"),
1474
+ high_count=sum(1 for v in all_vulnerabilities if v.severity == "High"),
1475
+ medium_count=sum(1 for v in all_vulnerabilities if v.severity == "Medium"),
1476
+ low_count=sum(1 for v in all_vulnerabilities if v.severity == "Low"),
1477
+ status="completed",
1478
+ )
1479
+
1480
+ db_session.add(report_record)
1481
+ await db_session.commit()
1482
+ await db_session.refresh(report_record)
1483
+
1484
+ logger.info(f"Successfully generated vulnerability report: {report_name}")
1485
+
1486
+ # Build response
1487
+ report_response = VulnerabilityReportResponse(
1488
+ id=report_record.id,
1489
+ report_name=report_record.report_name,
1490
+ customer_code=report_record.customer_code,
1491
+ file_name=report_record.file_name,
1492
+ file_size=report_record.file_size,
1493
+ generated_at=report_record.generated_at,
1494
+ generated_by=report_record.generated_by,
1495
+ total_vulnerabilities=report_record.total_vulnerabilities,
1496
+ critical_count=report_record.critical_count,
1497
+ high_count=report_record.high_count,
1498
+ medium_count=report_record.medium_count,
1499
+ low_count=report_record.low_count,
1500
+ filters_applied=json.loads(report_record.filters_json or "{}"),
1501
+ status=report_record.status,
1502
+ download_url=f"/api/v1/vulnerabilities/reports/{report_record.id}/download",
1503
+ )
1504
+
1505
+ return VulnerabilityReportGenerateResponse(
1506
+ success=True,
1507
+ message=f"Successfully generated report with {len(all_vulnerabilities)} vulnerabilities",
1508
+ report=report_response,
1509
+ )
1510
+
1511
+ except Exception as e:
1512
+ logger.error(f"Error generating vulnerability report: {e}")
1513
+
1514
+ # If we have a report_id, update it to failed status
1515
+ if report_id:
1516
+ try:
1517
+ stmt = select(VulnerabilityReport).filter(VulnerabilityReport.id == report_id)
1518
+ result = await db_session.execute(stmt)
1519
+ report = result.scalars().first()
1520
+ if report:
1521
+ report.status = "failed"
1522
+ report.error_message = str(e)
1523
+ await db_session.commit()
1524
+ except Exception as update_error:
1525
+ logger.error(f"Failed to update report status: {update_error}")
1526
+
1527
+ return VulnerabilityReportGenerateResponse(
1528
+ success=False,
1529
+ message="Failed to generate vulnerability report",
1530
+ error=str(e),
1531
+ )
1532
+
1533
+
1534
+async def list_vulnerability_reports(
1535
+ db_session: AsyncSession,
1536
+ current_user: User,
1537
+ customer_code: Optional[str] = None,
1538
+) -> VulnerabilityReportListResponse:
1539
+ """
1540
+ List available vulnerability reports
1541
+
1542
+ Args:
1543
+ db_session: Database session
1544
+ current_user: Current authenticated user
1545
+ customer_code: Optional filter by customer code
1546
+
1547
+ Returns:
1548
+ VulnerabilityReportListResponse with list of reports
1549
+ """
1550
+ try:
1551
+ # Get accessible customers
1552
+ accessible_customers = await customer_access_handler.get_user_accessible_customers(current_user, db_session)
1553
+
1554
+ # Build query
1555
+ query = select(VulnerabilityReport).order_by(desc(VulnerabilityReport.generated_at))
1556
+
1557
+ # Apply customer filtering
1558
+ if "*" not in accessible_customers:
1559
+ query = query.filter(VulnerabilityReport.customer_code.in_(accessible_customers))
1560
+
1561
+ if customer_code:
1562
+ if "*" not in accessible_customers and customer_code not in accessible_customers:
1563
+ return VulnerabilityReportListResponse(
1564
+ reports=[],
1565
+ total_count=0,
1566
+ success=True,
1567
+ message=f"Access denied to customer {customer_code}",
1568
+ )
1569
+ query = query.filter(VulnerabilityReport.customer_code == customer_code)
1570
+
1571
+ result = await db_session.execute(query)
1572
+ reports = result.scalars().all()
1573
+
1574
+ report_list = []
1575
+ for report in reports:
1576
+ report_response = VulnerabilityReportResponse(
1577
+ id=report.id,
1578
+ report_name=report.report_name,
1579
+ customer_code=report.customer_code,
1580
+ file_name=report.file_name,
1581
+ file_size=report.file_size,
1582
+ generated_at=report.generated_at,
1583
+ generated_by=report.generated_by,
1584
+ total_vulnerabilities=report.total_vulnerabilities,
1585
+ critical_count=report.critical_count,
1586
+ high_count=report.high_count,
1587
+ medium_count=report.medium_count,
1588
+ low_count=report.low_count,
1589
+ filters_applied=json.loads(report.filters_json or "{}"),
1590
+ status=report.status,
1591
+ download_url=f"/api/v1/vulnerabilities/reports/{report.id}/download",
1592
+ )
1593
+ report_list.append(report_response)
1594
+
1595
+ return VulnerabilityReportListResponse(
1596
+ reports=report_list,
1597
+ total_count=len(report_list),
1598
+ success=True,
1599
+ message=f"Found {len(report_list)} vulnerability reports",
1600
+ )
1601
+
1602
+ except Exception as e:
1603
+ logger.error(f"Error listing vulnerability reports: {e}")
1604
+ return VulnerabilityReportListResponse(
1605
+ reports=[],
1606
+ total_count=0,
1607
+ success=False,
1608
+ message=f"Failed to list reports: {e}",
1609
+ )
1610
+
1611
+
1612
+async def get_vulnerability_report_download(
1613
+ db_session: AsyncSession,
1614
+ current_user: User,
1615
+ report_id: int,
1616
+) -> Dict[str, Any]:
1617
+ """
1618
+ Get vulnerability report for download
1619
+
1620
+ Args:
1621
+ db_session: Database session
1622
+ current_user: Current authenticated user
1623
+ report_id: Report ID to download
1624
+
1625
+ Returns:
1626
+ Dict with file_content, file_name, and content_type
1627
+ """
1628
+ try:
1629
+ # Get report record
1630
+ result = await db_session.execute(select(VulnerabilityReport).filter(VulnerabilityReport.id == report_id))
1631
+ report = result.scalars().first()
1632
+
1633
+ if not report:
1634
+ raise HTTPException(status_code=404, detail="Report not found")
1635
+
1636
+ # Verify customer access
1637
+ accessible_customers = await customer_access_handler.get_user_accessible_customers(current_user, db_session)
1638
+
1639
+ if "*" not in accessible_customers and report.customer_code not in accessible_customers:
1640
+ raise HTTPException(status_code=403, detail="Access denied to this report")
1641
+
1642
+ # Retrieve file from MinIO
1643
+ from app.data_store.data_store_operations import retrieve_file_from_minio
1644
+
1645
+ file_data = await retrieve_file_from_minio(
1646
+ bucket_name=report.bucket_name,
1647
+ object_key=report.object_key,
1648
+ )
1649
+
1650
+ if not file_data["success"]:
1651
+ raise HTTPException(status_code=500, detail="Failed to retrieve report file")
1652
+
1653
+ return {
1654
+ "file_content": file_data["file_content"],
1655
+ "file_name": report.file_name,
1656
+ "content_type": "text/csv",
1657
+ }
1658
+
1659
+ except HTTPException:
1660
+ raise
1661
+ except Exception as e:
1662
+ logger.error(f"Error retrieving vulnerability report: {e}")
1663
+ raise HTTPException(status_code=500, detail=f"Failed to retrieve report: {e}")
1664
+
1665
+
1666
+async def delete_vulnerability_report(
1667
+ db_session: AsyncSession,
1668
+ current_user: User,
1669
+ report_id: int,
1670
+) -> Dict[str, Any]:
1671
+ """
1672
+ Delete a vulnerability report and its associated file from MinIO.
1673
+
1674
+ Args:
1675
+ db_session: Database session
1676
+ current_user: Current authenticated user
1677
+ report_id: ID of the report to delete
1678
+
1679
+ Returns:
1680
+ Dict with success status and details
1681
+ """
1682
+ from app.data_store.data_store_operations import delete_file_from_minio
1683
+
1684
+ try:
1685
+ # Get the report record
1686
+ stmt = select(VulnerabilityReport).filter(VulnerabilityReport.id == report_id)
1687
+ result = await db_session.execute(stmt)
1688
+ report = result.scalars().first()
1689
+
1690
+ if not report:
1691
+ return {
1692
+ "success": False,
1693
+ "error": f"Report with ID {report_id} not found",
1694
+ }
1695
+
1696
+ # Verify customer access
1697
+ accessible_customers = await customer_access_handler.get_user_accessible_customers(current_user, db_session)
1698
+
1699
+ if "*" not in accessible_customers and report.customer_code not in accessible_customers:
1700
+ return {
1701
+ "success": False,
1702
+ "error": f"Access denied to delete report for customer {report.customer_code}",
1703
+ }
1704
+
1705
+ logger.info(f"Deleting vulnerability report ID {report_id} for customer {report.customer_code}")
1706
+
1707
+ # Delete file from MinIO
1708
+ minio_result = await delete_file_from_minio(
1709
+ bucket_name=report.bucket_name,
1710
+ object_key=report.object_key,
1711
+ )
1712
+
1713
+ if not minio_result["success"]:
1714
+ logger.warning(
1715
+ f"Failed to delete file from MinIO for report {report_id}: {minio_result.get('error')}. "
1716
+ "Proceeding with database deletion.",
1717
+ )
1718
+
1719
+ # Store report details before deletion
1720
+ report_name = report.report_name
1721
+ customer_code = report.customer_code
1722
+
1723
+ # Delete database record
1724
+ await db_session.delete(report)
1725
+ await db_session.commit()
1726
+
1727
+ logger.info(f"Successfully deleted vulnerability report ID {report_id}")
1728
+
1729
+ return {
1730
+ "success": True,
1731
+ "message": f"Report '{report_name}' deleted successfully",
1732
+ "report_id": report_id,
1733
+ "report_name": report_name,
1734
+ "customer_code": customer_code,
1735
+ }
1736
+
1737
+ except Exception as e:
1738
+ logger.error(f"Error deleting vulnerability report {report_id}: {e}")
1739
+ return {
1740
+ "success": False,
1741
+ "error": str(e),
1742
+ }