main
py 2,001 lines 77.4 KB
Raw
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 }