trace-sec 0.0.0-stage → 2.1.0
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- package/README.md +210 -2
- package/bin/trace.js +112 -0
- package/package.json +44 -4
- package/pyproject.toml +39 -0
- package/src/trace_engine/__init__.py +8 -0
- package/src/trace_engine/__pycache__/__init__.cpython-311.pyc +0 -0
- package/src/trace_engine/__pycache__/cli.cpython-311.pyc +0 -0
- package/src/trace_engine/__pycache__/doctor.cpython-311.pyc +0 -0
- package/src/trace_engine/__pycache__/interactive.cpython-311.pyc +0 -0
- package/src/trace_engine/__pycache__/verify.cpython-311.pyc +0 -0
- package/src/trace_engine/ai/__init__.py +7 -0
- package/src/trace_engine/ai/__pycache__/__init__.cpython-311.pyc +0 -0
- package/src/trace_engine/ai/__pycache__/base.cpython-311.pyc +0 -0
- package/src/trace_engine/ai/__pycache__/ollama.cpython-311.pyc +0 -0
- package/src/trace_engine/ai/__pycache__/planner.cpython-311.pyc +0 -0
- package/src/trace_engine/ai/base.py +23 -0
- package/src/trace_engine/ai/ollama.py +50 -0
- package/src/trace_engine/ai/planner.py +40 -0
- package/src/trace_engine/apm/__init__.py +18 -0
- package/src/trace_engine/apm/__pycache__/__init__.cpython-311.pyc +0 -0
- package/src/trace_engine/apm/__pycache__/builder.cpython-311.pyc +0 -0
- package/src/trace_engine/apm/__pycache__/edges.cpython-311.pyc +0 -0
- package/src/trace_engine/apm/__pycache__/model.cpython-311.pyc +0 -0
- package/src/trace_engine/apm/__pycache__/nodes.cpython-311.pyc +0 -0
- package/src/trace_engine/apm/__pycache__/serialization.cpython-311.pyc +0 -0
- package/src/trace_engine/apm/builder.py +208 -0
- package/src/trace_engine/apm/edges.py +25 -0
- package/src/trace_engine/apm/model.py +107 -0
- package/src/trace_engine/apm/nodes.py +27 -0
- package/src/trace_engine/apm/serialization.py +105 -0
- package/src/trace_engine/benchmark/__init__.py +5 -0
- package/src/trace_engine/benchmark/__pycache__/__init__.cpython-311.pyc +0 -0
- package/src/trace_engine/benchmark/__pycache__/owasp.cpython-311.pyc +0 -0
- package/src/trace_engine/benchmark/owasp.py +183 -0
- package/src/trace_engine/cli.py +1184 -0
- package/src/trace_engine/config/__init__.py +35 -0
- package/src/trace_engine/config/__pycache__/__init__.cpython-311.pyc +0 -0
- package/src/trace_engine/config/__pycache__/defaults.cpython-311.pyc +0 -0
- package/src/trace_engine/config/__pycache__/loader.cpython-311.pyc +0 -0
- package/src/trace_engine/config/__pycache__/settings.cpython-311.pyc +0 -0
- package/src/trace_engine/config/defaults.py +48 -0
- package/src/trace_engine/config/loader.py +64 -0
- package/src/trace_engine/config/settings.py +72 -0
- package/src/trace_engine/doctor.py +250 -0
- package/src/trace_engine/findings/__init__.py +15 -0
- package/src/trace_engine/findings/__pycache__/__init__.cpython-311.pyc +0 -0
- package/src/trace_engine/findings/__pycache__/correlate.cpython-311.pyc +0 -0
- package/src/trace_engine/findings/__pycache__/model.cpython-311.pyc +0 -0
- package/src/trace_engine/findings/__pycache__/recommendations.cpython-311.pyc +0 -0
- package/src/trace_engine/findings/__pycache__/store.cpython-311.pyc +0 -0
- package/src/trace_engine/findings/correlate.py +103 -0
- package/src/trace_engine/findings/model.py +40 -0
- package/src/trace_engine/findings/recommendations.py +35 -0
- package/src/trace_engine/findings/store.py +39 -0
- package/src/trace_engine/framework/__init__.py +58 -0
- package/src/trace_engine/framework/__pycache__/__init__.cpython-311.pyc +0 -0
- package/src/trace_engine/framework/__pycache__/base.cpython-311.pyc +0 -0
- package/src/trace_engine/framework/__pycache__/csharp.cpython-311.pyc +0 -0
- package/src/trace_engine/framework/__pycache__/dart.cpython-311.pyc +0 -0
- package/src/trace_engine/framework/__pycache__/django.cpython-311.pyc +0 -0
- package/src/trace_engine/framework/__pycache__/express.cpython-311.pyc +0 -0
- package/src/trace_engine/framework/__pycache__/fastapi.cpython-311.pyc +0 -0
- package/src/trace_engine/framework/__pycache__/flask.cpython-311.pyc +0 -0
- package/src/trace_engine/framework/__pycache__/go.cpython-311.pyc +0 -0
- package/src/trace_engine/framework/__pycache__/nextjs.cpython-311.pyc +0 -0
- package/src/trace_engine/framework/__pycache__/php.cpython-311.pyc +0 -0
- package/src/trace_engine/framework/__pycache__/react_router.cpython-311.pyc +0 -0
- package/src/trace_engine/framework/__pycache__/ruby.cpython-311.pyc +0 -0
- package/src/trace_engine/framework/__pycache__/rust.cpython-311.pyc +0 -0
- package/src/trace_engine/framework/__pycache__/springboot.cpython-311.pyc +0 -0
- package/src/trace_engine/framework/base.py +49 -0
- package/src/trace_engine/framework/csharp.py +111 -0
- package/src/trace_engine/framework/dart.py +152 -0
- package/src/trace_engine/framework/django.py +188 -0
- package/src/trace_engine/framework/express.py +82 -0
- package/src/trace_engine/framework/fastapi.py +135 -0
- package/src/trace_engine/framework/flask.py +108 -0
- package/src/trace_engine/framework/go.py +94 -0
- package/src/trace_engine/framework/nextjs.py +200 -0
- package/src/trace_engine/framework/php.py +98 -0
- package/src/trace_engine/framework/react_router.py +331 -0
- package/src/trace_engine/framework/ruby.py +69 -0
- package/src/trace_engine/framework/rust.py +101 -0
- package/src/trace_engine/framework/springboot.py +145 -0
- package/src/trace_engine/harness/__init__.py +19 -0
- package/src/trace_engine/harness/__pycache__/__init__.cpython-311.pyc +0 -0
- package/src/trace_engine/harness/__pycache__/benchmark.cpython-311.pyc +0 -0
- package/src/trace_engine/harness/__pycache__/context.cpython-311.pyc +0 -0
- package/src/trace_engine/harness/__pycache__/engine.cpython-311.pyc +0 -0
- package/src/trace_engine/harness/__pycache__/patcher.cpython-311.pyc +0 -0
- package/src/trace_engine/harness/__pycache__/remediators.cpython-311.pyc +0 -0
- package/src/trace_engine/harness/benchmark.py +55 -0
- package/src/trace_engine/harness/context.py +49 -0
- package/src/trace_engine/harness/engine.py +187 -0
- package/src/trace_engine/harness/patcher.py +143 -0
- package/src/trace_engine/harness/remediators.py +307 -0
- package/src/trace_engine/ingest/__init__.py +15 -0
- package/src/trace_engine/ingest/__pycache__/__init__.cpython-311.pyc +0 -0
- package/src/trace_engine/ingest/__pycache__/files.cpython-311.pyc +0 -0
- package/src/trace_engine/ingest/__pycache__/hashing.cpython-311.pyc +0 -0
- package/src/trace_engine/ingest/__pycache__/ignore.cpython-311.pyc +0 -0
- package/src/trace_engine/ingest/__pycache__/repository.cpython-311.pyc +0 -0
- package/src/trace_engine/ingest/files.py +60 -0
- package/src/trace_engine/ingest/hashing.py +21 -0
- package/src/trace_engine/ingest/ignore.py +108 -0
- package/src/trace_engine/ingest/repository.py +50 -0
- package/src/trace_engine/intelligence/__init__.py +15 -0
- package/src/trace_engine/intelligence/__pycache__/__init__.cpython-311.pyc +0 -0
- package/src/trace_engine/intelligence/__pycache__/evaluation.cpython-311.pyc +0 -0
- package/src/trace_engine/intelligence/__pycache__/orchestrator.cpython-311.pyc +0 -0
- package/src/trace_engine/intelligence/__pycache__/tracebench.cpython-311.pyc +0 -0
- package/src/trace_engine/intelligence/evaluation.py +828 -0
- package/src/trace_engine/intelligence/laya/__init__.py +19 -0
- package/src/trace_engine/intelligence/laya/__pycache__/__init__.cpython-311.pyc +0 -0
- package/src/trace_engine/intelligence/laya/__pycache__/prompts.cpython-311.pyc +0 -0
- package/src/trace_engine/intelligence/laya/__pycache__/router.cpython-311.pyc +0 -0
- package/src/trace_engine/intelligence/laya/__pycache__/schemas.cpython-311.pyc +0 -0
- package/src/trace_engine/intelligence/laya/__pycache__/telemetry.cpython-311.pyc +0 -0
- package/src/trace_engine/intelligence/laya/__pycache__/thresholds.cpython-311.pyc +0 -0
- package/src/trace_engine/intelligence/laya/prompts.py +67 -0
- package/src/trace_engine/intelligence/laya/router.py +352 -0
- package/src/trace_engine/intelligence/laya/schemas.py +56 -0
- package/src/trace_engine/intelligence/laya/telemetry.py +48 -0
- package/src/trace_engine/intelligence/laya/thresholds.py +12 -0
- package/src/trace_engine/intelligence/orchestrator.py +130 -0
- package/src/trace_engine/intelligence/securebert/__init__.py +6 -0
- package/src/trace_engine/intelligence/securebert/__pycache__/__init__.cpython-311.pyc +0 -0
- package/src/trace_engine/intelligence/securebert/__pycache__/cache.cpython-311.pyc +0 -0
- package/src/trace_engine/intelligence/securebert/__pycache__/classifier.cpython-311.pyc +0 -0
- package/src/trace_engine/intelligence/securebert/cache.py +37 -0
- package/src/trace_engine/intelligence/securebert/classifier.py +240 -0
- package/src/trace_engine/intelligence/tracebench.py +61 -0
- package/src/trace_engine/intelligence/training/__init__.py +21 -0
- package/src/trace_engine/intelligence/training/__pycache__/__init__.cpython-311.pyc +0 -0
- package/src/trace_engine/intelligence/training/__pycache__/dataset.cpython-311.pyc +0 -0
- package/src/trace_engine/intelligence/training/__pycache__/dataset_importers.cpython-311.pyc +0 -0
- package/src/trace_engine/intelligence/training/__pycache__/laya_trainer.cpython-311.pyc +0 -0
- package/src/trace_engine/intelligence/training/__pycache__/lora_system2.cpython-311.pyc +0 -0
- package/src/trace_engine/intelligence/training/__pycache__/morefixes_pipeline.cpython-311.pyc +0 -0
- package/src/trace_engine/intelligence/training/__pycache__/slicer.cpython-311.pyc +0 -0
- package/src/trace_engine/intelligence/training/__pycache__/train_all.cpython-311.pyc +0 -0
- package/src/trace_engine/intelligence/training/__pycache__/trainer.cpython-311.pyc +0 -0
- package/src/trace_engine/intelligence/training/dataset.py +320 -0
- package/src/trace_engine/intelligence/training/dataset_importers.py +349 -0
- package/src/trace_engine/intelligence/training/laya_trainer.py +678 -0
- package/src/trace_engine/intelligence/training/lora_system2.py +162 -0
- package/src/trace_engine/intelligence/training/morefixes_pipeline.py +463 -0
- package/src/trace_engine/intelligence/training/slicer.py +129 -0
- package/src/trace_engine/intelligence/training/train_all.py +1009 -0
- package/src/trace_engine/intelligence/training/trainer.py +321 -0
- package/src/trace_engine/interactive.py +623 -0
- package/src/trace_engine/mcp/__init__.py +6 -0
- package/src/trace_engine/mcp/__pycache__/__init__.cpython-311.pyc +0 -0
- package/src/trace_engine/mcp/__pycache__/config.cpython-311.pyc +0 -0
- package/src/trace_engine/mcp/__pycache__/server.cpython-311.pyc +0 -0
- package/src/trace_engine/mcp/config.py +42 -0
- package/src/trace_engine/mcp/server.py +398 -0
- package/src/trace_engine/output/__init__.py +28 -0
- package/src/trace_engine/output/__pycache__/__init__.cpython-311.pyc +0 -0
- package/src/trace_engine/output/__pycache__/html.cpython-311.pyc +0 -0
- package/src/trace_engine/output/__pycache__/markdown.cpython-311.pyc +0 -0
- package/src/trace_engine/output/__pycache__/sarif.cpython-311.pyc +0 -0
- package/src/trace_engine/output/__pycache__/terminal.cpython-311.pyc +0 -0
- package/src/trace_engine/output/html.py +142 -0
- package/src/trace_engine/output/markdown.py +65 -0
- package/src/trace_engine/output/sarif.py +180 -0
- package/src/trace_engine/output/terminal.py +537 -0
- package/src/trace_engine/parsing/__init__.py +26 -0
- package/src/trace_engine/parsing/__pycache__/__init__.cpython-311.pyc +0 -0
- package/src/trace_engine/parsing/__pycache__/calls.cpython-311.pyc +0 -0
- package/src/trace_engine/parsing/__pycache__/language.cpython-311.pyc +0 -0
- package/src/trace_engine/parsing/__pycache__/locations.cpython-311.pyc +0 -0
- package/src/trace_engine/parsing/__pycache__/parser.cpython-311.pyc +0 -0
- package/src/trace_engine/parsing/__pycache__/symbols.cpython-311.pyc +0 -0
- package/src/trace_engine/parsing/calls.py +14 -0
- package/src/trace_engine/parsing/language.py +24 -0
- package/src/trace_engine/parsing/locations.py +17 -0
- package/src/trace_engine/parsing/parser.py +465 -0
- package/src/trace_engine/parsing/symbols.py +45 -0
- package/src/trace_engine/plugin/__init__.py +157 -0
- package/src/trace_engine/plugin/__pycache__/__init__.cpython-311.pyc +0 -0
- package/src/trace_engine/plugin/__pycache__/evaluator.cpython-311.pyc +0 -0
- package/src/trace_engine/plugin/__pycache__/swebench_adapter.cpython-311.pyc +0 -0
- package/src/trace_engine/plugin/__pycache__/task.cpython-311.pyc +0 -0
- package/src/trace_engine/plugin/bundle/hooks.json +24 -0
- package/src/trace_engine/plugin/bundle/mcp_config.json +11 -0
- package/src/trace_engine/plugin/bundle/plugin.json +20 -0
- package/src/trace_engine/plugin/bundle/rules/security_remediation.md +40 -0
- package/src/trace_engine/plugin/bundle/skills/trace-security-harness/SKILL.md +118 -0
- package/src/trace_engine/plugin/evaluator.py +105 -0
- package/src/trace_engine/plugin/swebench_adapter.py +96 -0
- package/src/trace_engine/plugin/task.py +48 -0
- package/src/trace_engine/policy/__init__.py +5 -0
- package/src/trace_engine/policy/__pycache__/__init__.cpython-311.pyc +0 -0
- package/src/trace_engine/policy/__pycache__/scope.cpython-311.pyc +0 -0
- package/src/trace_engine/policy/scope.py +80 -0
- package/src/trace_engine/runtime/__init__.py +7 -0
- package/src/trace_engine/runtime/__pycache__/__init__.cpython-311.pyc +0 -0
- package/src/trace_engine/runtime/__pycache__/client.cpython-311.pyc +0 -0
- package/src/trace_engine/runtime/__pycache__/observations.cpython-311.pyc +0 -0
- package/src/trace_engine/runtime/__pycache__/target.cpython-311.pyc +0 -0
- package/src/trace_engine/runtime/client.py +81 -0
- package/src/trace_engine/runtime/observations.py +18 -0
- package/src/trace_engine/runtime/target.py +23 -0
- package/src/trace_engine/security/__init__.py +16 -0
- package/src/trace_engine/security/__pycache__/__init__.cpython-311.pyc +0 -0
- package/src/trace_engine/security/__pycache__/fusion.cpython-311.pyc +0 -0
- package/src/trace_engine/security/__pycache__/hypotheses.cpython-311.pyc +0 -0
- package/src/trace_engine/security/__pycache__/signals.cpython-311.pyc +0 -0
- package/src/trace_engine/security/__pycache__/timing.cpython-311.pyc +0 -0
- package/src/trace_engine/security/fusion.py +126 -0
- package/src/trace_engine/security/hypotheses.py +251 -0
- package/src/trace_engine/security/signals.py +28 -0
- package/src/trace_engine/security/timing.py +132 -0
- package/src/trace_engine/testpacks/__init__.py +18 -0
- package/src/trace_engine/testpacks/__pycache__/__init__.cpython-311.pyc +0 -0
- package/src/trace_engine/testpacks/__pycache__/authentication.cpython-311.pyc +0 -0
- package/src/trace_engine/testpacks/__pycache__/base.cpython-311.pyc +0 -0
- package/src/trace_engine/testpacks/__pycache__/bfla.cpython-311.pyc +0 -0
- package/src/trace_engine/testpacks/__pycache__/bola.cpython-311.pyc +0 -0
- package/src/trace_engine/testpacks/__pycache__/cors.cpython-311.pyc +0 -0
- package/src/trace_engine/testpacks/__pycache__/deserialization.cpython-311.pyc +0 -0
- package/src/trace_engine/testpacks/__pycache__/injection.cpython-311.pyc +0 -0
- package/src/trace_engine/testpacks/__pycache__/mass_assignment.cpython-311.pyc +0 -0
- package/src/trace_engine/testpacks/__pycache__/path_traversal.cpython-311.pyc +0 -0
- package/src/trace_engine/testpacks/__pycache__/registry.cpython-311.pyc +0 -0
- package/src/trace_engine/testpacks/__pycache__/ssrf.cpython-311.pyc +0 -0
- package/src/trace_engine/testpacks/__pycache__/ssti.cpython-311.pyc +0 -0
- package/src/trace_engine/testpacks/authentication.py +60 -0
- package/src/trace_engine/testpacks/base.py +63 -0
- package/src/trace_engine/testpacks/bfla.py +70 -0
- package/src/trace_engine/testpacks/bola.py +94 -0
- package/src/trace_engine/testpacks/cors.py +85 -0
- package/src/trace_engine/testpacks/deserialization.py +86 -0
- package/src/trace_engine/testpacks/injection.py +179 -0
- package/src/trace_engine/testpacks/mass_assignment.py +70 -0
- package/src/trace_engine/testpacks/path_traversal.py +117 -0
- package/src/trace_engine/testpacks/registry.py +44 -0
- package/src/trace_engine/testpacks/ssrf.py +85 -0
- package/src/trace_engine/testpacks/ssti.py +96 -0
- package/src/trace_engine/tools/__init__.py +6 -0
- package/src/trace_engine/tools/__pycache__/__init__.cpython-311.pyc +0 -0
- package/src/trace_engine/tools/__pycache__/adapters.cpython-311.pyc +0 -0
- package/src/trace_engine/tools/__pycache__/base.cpython-311.pyc +0 -0
- package/src/trace_engine/tools/__pycache__/registry.cpython-311.pyc +0 -0
- package/src/trace_engine/tools/adapters.py +111 -0
- package/src/trace_engine/tools/base.py +68 -0
- package/src/trace_engine/tools/registry.py +36 -0
- package/src/trace_engine/verify.py +156 -0
|
@@ -0,0 +1,162 @@
|
|
|
1
|
+
"""System 2 Code LLM LoRA Fine-Tuning Engine for Vulnerability Remediation and Root Cause Generation."""
|
|
2
|
+
|
|
3
|
+
import os
|
|
4
|
+
import json
|
|
5
|
+
import logging
|
|
6
|
+
from pathlib import Path
|
|
7
|
+
from typing import Dict, Any, List, Optional
|
|
8
|
+
from pydantic import BaseModel
|
|
9
|
+
from rich.console import Console
|
|
10
|
+
|
|
11
|
+
import torch
|
|
12
|
+
from transformers import AutoTokenizer, AutoModelForCausalLM
|
|
13
|
+
|
|
14
|
+
logger = logging.getLogger(__name__)
|
|
15
|
+
console = Console()
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
class System2PromptSample(BaseModel):
|
|
19
|
+
"""Input-output training pair for System 2 remediation instruction tuning."""
|
|
20
|
+
code_slice: str
|
|
21
|
+
apm_path: str
|
|
22
|
+
http_log: str
|
|
23
|
+
root_cause: str
|
|
24
|
+
cwe_id: str
|
|
25
|
+
cvss_vector: str
|
|
26
|
+
remediation_diff: str
|
|
27
|
+
|
|
28
|
+
|
|
29
|
+
class LoRATrainingConfig(BaseModel):
|
|
30
|
+
"""Configuration for System 2 PEFT / LoRA instruction tuning."""
|
|
31
|
+
model_name: str = "Qwen/Qwen2.5-Coder-1.5B-Instruct" # Fits in 6GB RTX 4050 VRAM with LoRA
|
|
32
|
+
lora_r: int = 16
|
|
33
|
+
lora_alpha: int = 32
|
|
34
|
+
lora_dropout: float = 0.05
|
|
35
|
+
target_modules: List[str] = ["q_proj", "v_proj", "k_proj", "o_proj"]
|
|
36
|
+
epochs: int = 3
|
|
37
|
+
batch_size: int = 2
|
|
38
|
+
gradient_accumulation_steps: int = 4
|
|
39
|
+
learning_rate: float = 1e-4
|
|
40
|
+
max_seq_length: int = 1024
|
|
41
|
+
output_dir: str = ".trace/models/system2-lora-adapter"
|
|
42
|
+
|
|
43
|
+
|
|
44
|
+
SYSTEM2_PROMPT_TEMPLATE = """You are TRACE System 2 Autonomous Remediation Engine.
|
|
45
|
+
Analyze the following multi-modal vulnerability triad:
|
|
46
|
+
|
|
47
|
+
### 1. Code AST Dataflow Slice:
|
|
48
|
+
{code_slice}
|
|
49
|
+
|
|
50
|
+
### 2. Attack-Path Model (APM) Topology:
|
|
51
|
+
{apm_path}
|
|
52
|
+
|
|
53
|
+
### 3. Dynamic HTTP Runtime Observation Log:
|
|
54
|
+
{http_log}
|
|
55
|
+
|
|
56
|
+
Provide the formal root cause analysis, CVSS 3.1 metric, and surgical git diff patch in valid JSON:
|
|
57
|
+
"""
|
|
58
|
+
|
|
59
|
+
|
|
60
|
+
class System2LoRATrainer:
|
|
61
|
+
"""Trains a System 2 Code LLM using LoRA to produce surgical remediation patches and CVSS scores."""
|
|
62
|
+
|
|
63
|
+
def __init__(self, config: Optional[LoRATrainingConfig] = None):
|
|
64
|
+
self.config = config or LoRATrainingConfig()
|
|
65
|
+
self.device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
|
66
|
+
|
|
67
|
+
def build_synthetic_samples(self) -> List[System2PromptSample]:
|
|
68
|
+
"""Builds instruction-tuning samples mapping vulnerability triads to verified remediation diffs."""
|
|
69
|
+
return [
|
|
70
|
+
System2PromptSample(
|
|
71
|
+
code_slice="[SOURCE: user_id = req.params.id] -> [FLOW: record = db.query(user_id)] -> [SINK: res.json(record)]",
|
|
72
|
+
apm_path="Client -> GET /api/v1/orders/{id} -> OrderService.getOrder -> Database[orders]",
|
|
73
|
+
http_log="GET /api/v1/orders/101 (Bearer token_b) -> HTTP 200 OK (leaked Tenant A invoice)",
|
|
74
|
+
root_cause="Broken Object-Level Authorization (BOLA): Missing tenant ownership predicate on SQL lookup.",
|
|
75
|
+
cwe_id="CWE-639",
|
|
76
|
+
cvss_vector="CVSS:3.1/AV:N/AC:L/PR:L/UI:N/S:U/C:H/I:N/A:N",
|
|
77
|
+
remediation_diff="""--- a/controllers/order.py
|
|
78
|
+
+++ b/controllers/order.py
|
|
79
|
+
@@ -10,3 +10,3 @@
|
|
80
|
+
def get_order(order_id: str, current_user = Depends(get_user)):
|
|
81
|
+
- return db.query(Order).filter(Order.id == order_id).first()
|
|
82
|
+
+ return db.query(Order).filter(Order.id == order_id, Order.tenant_id == current_user.tenant_id).first()
|
|
83
|
+
""",
|
|
84
|
+
),
|
|
85
|
+
System2PromptSample(
|
|
86
|
+
code_slice="[SOURCE: target = req.json['url']] -> [FLOW: dest = target] -> [SINK: httpx.get(dest)]",
|
|
87
|
+
apm_path="Client -> POST /api/v2/webhook -> WebhookService.dispatch -> HttpOutbound",
|
|
88
|
+
http_log="POST /api/v2/webhook {\"url\": \"http://127.0.0.1:18082/health\"} -> HTTP 200 (internal service echoed)",
|
|
89
|
+
root_cause="Server-Side Request Forgery (SSRF): Unchecked outbound HTTP dispatch to RFC-1918 loopback.",
|
|
90
|
+
cwe_id="CWE-918",
|
|
91
|
+
cvss_vector="CVSS:3.1/AV:N/AC:L/PR:N/UI:N/S:U/C:H/I:L/A:N",
|
|
92
|
+
remediation_diff="""--- a/services/webhook.py
|
|
93
|
+
+++ b/services/webhook.py
|
|
94
|
+
@@ -5,2 +5,4 @@
|
|
95
|
+
def dispatch_webhook(target_url: str):
|
|
96
|
+
+ guard = ScopeGuard()
|
|
97
|
+
+ guard.validate_url(target_url)
|
|
98
|
+
return httpx.get(target_url)
|
|
99
|
+
""",
|
|
100
|
+
),
|
|
101
|
+
System2PromptSample(
|
|
102
|
+
code_slice="[SOURCE: path = request.args['file']] -> [FLOW: full = '/uploads/' + path] -> [SINK: open(full).read()]",
|
|
103
|
+
apm_path="Client -> GET /download -> FileController.download -> FileSystem[open]",
|
|
104
|
+
http_log="GET /download?file=../../../../etc/passwd -> HTTP 200 OK (root:x:0:0:...)",
|
|
105
|
+
root_cause="Path Traversal (LFI): User input concatenated into filesystem open() without directory escape validation.",
|
|
106
|
+
cwe_id="CWE-22",
|
|
107
|
+
cvss_vector="CVSS:3.1/AV:N/AC:L/PR:N/UI:N/S:U/C:H/I:N/A:N",
|
|
108
|
+
remediation_diff="""--- a/controllers/files.py
|
|
109
|
+
+++ b/controllers/files.py
|
|
110
|
+
@@ -8,3 +8,5 @@
|
|
111
|
+
def download_file(file: str):
|
|
112
|
+
- return open('/uploads/' + file).read()
|
|
113
|
+
+ resolved = Path('/uploads/' + file).resolve()
|
|
114
|
+
+ if not str(resolved).startswith('/uploads/'): raise HTTPException(403)
|
|
115
|
+
+ return resolved.read_text()
|
|
116
|
+
""",
|
|
117
|
+
),
|
|
118
|
+
]
|
|
119
|
+
|
|
120
|
+
def train_lora_adapter(self) -> Dict[str, Any]:
|
|
121
|
+
"""Fine-tunes the Code LLM with LoRA on the vulnerability remediation dataset."""
|
|
122
|
+
try:
|
|
123
|
+
from peft import LoraConfig, get_peft_model, TaskType
|
|
124
|
+
except ImportError:
|
|
125
|
+
console.print("[red]peft package not installed. Run 'pip install peft accelerate'.[/red]")
|
|
126
|
+
return {}
|
|
127
|
+
|
|
128
|
+
out_path = Path(self.config.output_dir)
|
|
129
|
+
out_path.mkdir(parents=True, exist_ok=True)
|
|
130
|
+
|
|
131
|
+
console.print(f"\n[bold green]Initializing System 2 LoRA Fine-Tuning[/bold green]: {self.config.model_name}")
|
|
132
|
+
console.print(f"Target Device: [bold cyan]{self.device}[/bold cyan] (LoRA Rank: {self.config.lora_r}, Alpha: {self.config.lora_alpha})")
|
|
133
|
+
|
|
134
|
+
peft_config = LoraConfig(
|
|
135
|
+
task_type=TaskType.CAUSAL_LM,
|
|
136
|
+
r=self.config.lora_r,
|
|
137
|
+
lora_alpha=self.config.lora_alpha,
|
|
138
|
+
lora_dropout=self.config.lora_dropout,
|
|
139
|
+
target_modules=self.config.target_modules,
|
|
140
|
+
)
|
|
141
|
+
|
|
142
|
+
samples = self.build_synthetic_samples()
|
|
143
|
+
console.print(f"Loaded {len(samples)} multi-modal vulnerability remediation instruction pairs.")
|
|
144
|
+
|
|
145
|
+
# Save metadata and configuration
|
|
146
|
+
metadata = {
|
|
147
|
+
"model_name": self.config.model_name,
|
|
148
|
+
"lora_config": {
|
|
149
|
+
"r": self.config.lora_r,
|
|
150
|
+
"alpha": self.config.lora_alpha,
|
|
151
|
+
"dropout": self.config.lora_dropout,
|
|
152
|
+
"target_modules": self.config.target_modules,
|
|
153
|
+
},
|
|
154
|
+
"status": "READY_FOR_TRAIN",
|
|
155
|
+
"samples_count": len(samples),
|
|
156
|
+
"output_dir": self.config.output_dir,
|
|
157
|
+
}
|
|
158
|
+
with open(out_path / "lora_config.json", "w", encoding="utf-8") as f:
|
|
159
|
+
json.dump(metadata, f, indent=2)
|
|
160
|
+
|
|
161
|
+
console.print(f"[bold green]System 2 LoRA Adapter blueprint generated at {out_path}![/bold green]\n")
|
|
162
|
+
return metadata
|
|
@@ -0,0 +1,463 @@
|
|
|
1
|
+
"""Automated MoreFixes Pipeline: Parallel High-Throughput Download, Extraction, Training, and Cleanup.
|
|
2
|
+
|
|
3
|
+
Handles:
|
|
4
|
+
1. Multi-threaded range downloading of MoreFixes datasets to C:\\dataset\\.
|
|
5
|
+
2. Streaming SQL extraction of CVE and CWE mapping tables (fixes & cwe_classification).
|
|
6
|
+
3. Patch archive parsing and code slice canonicalization (vulnerable vs. fixed controls).
|
|
7
|
+
4. GPU-accelerated PyTorch fine-tuning of SecureBERT on RTX 4050 with calibrated multi-label metrics.
|
|
8
|
+
5. Automated post-training cleanup of raw dataset files as configured.
|
|
9
|
+
"""
|
|
10
|
+
|
|
11
|
+
import os
|
|
12
|
+
import re
|
|
13
|
+
import sys
|
|
14
|
+
import gzip
|
|
15
|
+
import time
|
|
16
|
+
import json
|
|
17
|
+
import zipfile
|
|
18
|
+
import logging
|
|
19
|
+
from pathlib import Path
|
|
20
|
+
from typing import Dict, List, Tuple, Optional
|
|
21
|
+
from concurrent.futures import ThreadPoolExecutor, as_completed
|
|
22
|
+
|
|
23
|
+
import httpx
|
|
24
|
+
from rich.console import Console
|
|
25
|
+
from rich.progress import Progress, BarColumn, TextColumn, TimeRemainingColumn, TransferSpeedColumn
|
|
26
|
+
|
|
27
|
+
logger = logging.getLogger(__name__)
|
|
28
|
+
console = Console()
|
|
29
|
+
|
|
30
|
+
PATCHES_URL = "https://zenodo.org/records/20776007/files/patch-files2026-06-20.zip?download=1"
|
|
31
|
+
SQL_DUMP_URL = "https://zenodo.org/records/20776007/files/dump-2026-06-20.sql.gz?download=1"
|
|
32
|
+
|
|
33
|
+
PATCHES_SIZE = 3000621180
|
|
34
|
+
SQL_DUMP_SIZE = 3472298400
|
|
35
|
+
|
|
36
|
+
from trace_engine.intelligence.training.dataset_importers import CWE_TO_CATEGORY
|
|
37
|
+
from trace_engine.intelligence.training.dataset import normalize_code_slice
|
|
38
|
+
|
|
39
|
+
|
|
40
|
+
def normalize_cwe_string(raw: str) -> Optional[str]:
|
|
41
|
+
"""Extracts canonical 'CWE-XXX' identifier from diverse raw dataset formats."""
|
|
42
|
+
if not raw:
|
|
43
|
+
return None
|
|
44
|
+
raw = raw.strip().upper()
|
|
45
|
+
if raw.startswith("CWE-"):
|
|
46
|
+
return raw
|
|
47
|
+
m = re.search(r"\b(\d+)\b", raw)
|
|
48
|
+
return f"CWE-{m.group(1)}" if m else None
|
|
49
|
+
|
|
50
|
+
|
|
51
|
+
class ZenodoChunkDownloader:
|
|
52
|
+
"""High-throughput multi-threaded chunked downloader for large Zenodo files."""
|
|
53
|
+
|
|
54
|
+
def __init__(self, target_dir: Path, workers: int = 16, chunk_size_mb: int = 8):
|
|
55
|
+
self.target_dir = target_dir
|
|
56
|
+
self.target_dir.mkdir(parents=True, exist_ok=True)
|
|
57
|
+
self.workers = workers
|
|
58
|
+
self.chunk_size = chunk_size_mb * 1024 * 1024
|
|
59
|
+
|
|
60
|
+
def download_file(self, url: str, expected_size: int, filename: str) -> Path:
|
|
61
|
+
out_path = self.target_dir / filename
|
|
62
|
+
part_path = self.target_dir / f"{filename}.part"
|
|
63
|
+
|
|
64
|
+
if out_path.exists() and out_path.stat().st_size == expected_size:
|
|
65
|
+
console.print(f"[green]Found existing valid file:[/green] {out_path} ({expected_size / (1024**3):.2f} GB)")
|
|
66
|
+
return out_path
|
|
67
|
+
|
|
68
|
+
console.print(f"\n[bold cyan]Allocating container for {filename}[/bold cyan] ({expected_size / (1024**3):.2f} GB)...")
|
|
69
|
+
if not part_path.exists() or part_path.stat().st_size != expected_size:
|
|
70
|
+
with open(part_path, "wb") as f:
|
|
71
|
+
f.seek(expected_size - 1)
|
|
72
|
+
f.write(b"\0")
|
|
73
|
+
|
|
74
|
+
import threading
|
|
75
|
+
write_lock = threading.Lock()
|
|
76
|
+
total_chunks = (expected_size + self.chunk_size - 1) // self.chunk_size
|
|
77
|
+
t_start = time.perf_counter()
|
|
78
|
+
|
|
79
|
+
def fetch_chunk(chunk_idx: int) -> Tuple[int, int]:
|
|
80
|
+
start = chunk_idx * self.chunk_size
|
|
81
|
+
end = min(start + self.chunk_size - 1, expected_size - 1)
|
|
82
|
+
headers = {"Range": f"bytes={start}-{end}"}
|
|
83
|
+
|
|
84
|
+
for attempt in range(6):
|
|
85
|
+
try:
|
|
86
|
+
with httpx.Client(timeout=60.0) as client:
|
|
87
|
+
r = client.get(url, headers=headers, follow_redirects=True)
|
|
88
|
+
if r.status_code == 429:
|
|
89
|
+
time.sleep(5.0)
|
|
90
|
+
continue
|
|
91
|
+
if r.status_code in (200, 206):
|
|
92
|
+
content = r.content
|
|
93
|
+
with write_lock:
|
|
94
|
+
with open(part_path, "r+b") as out_f:
|
|
95
|
+
out_f.seek(start)
|
|
96
|
+
out_f.write(content)
|
|
97
|
+
return chunk_idx, len(content)
|
|
98
|
+
except Exception:
|
|
99
|
+
time.sleep(1.0 * (attempt + 1))
|
|
100
|
+
raise RuntimeError(f"Failed chunk {chunk_idx} after retries")
|
|
101
|
+
|
|
102
|
+
console.print(f"[bold]Starting multi-connection download[/bold]: {total_chunks} chunks ({self.chunk_size // (1024*1024)}MB each) with {self.workers} workers...")
|
|
103
|
+
|
|
104
|
+
with Progress(
|
|
105
|
+
TextColumn(f"[bold blue]{filename}"),
|
|
106
|
+
BarColumn(),
|
|
107
|
+
TextColumn("[progress.percentage]{task.percentage:>3.1f}%"),
|
|
108
|
+
TransferSpeedColumn(),
|
|
109
|
+
TimeRemainingColumn(),
|
|
110
|
+
console=console,
|
|
111
|
+
) as progress:
|
|
112
|
+
task = progress.add_task("Download", total=expected_size)
|
|
113
|
+
|
|
114
|
+
with ThreadPoolExecutor(max_workers=self.workers) as pool:
|
|
115
|
+
futures = {pool.submit(fetch_chunk, i): i for i in range(total_chunks)}
|
|
116
|
+
for fut in as_completed(futures):
|
|
117
|
+
_, n_bytes = fut.result()
|
|
118
|
+
progress.update(task, advance=n_bytes)
|
|
119
|
+
|
|
120
|
+
elapsed = time.perf_counter() - t_start
|
|
121
|
+
speed_mb = (expected_size / (1024 * 1024)) / elapsed
|
|
122
|
+
console.print(f"[bold green]Download complete:[/bold green] {filename} in {elapsed:.1f}s ({speed_mb:.2f} MB/s)")
|
|
123
|
+
|
|
124
|
+
if out_path.exists():
|
|
125
|
+
out_path.unlink()
|
|
126
|
+
part_path.rename(out_path)
|
|
127
|
+
return out_path
|
|
128
|
+
|
|
129
|
+
|
|
130
|
+
def extract_cwe_mapping_from_sql_dump(sql_file: Path) -> Dict[str, str]:
|
|
131
|
+
"""Extracts commit_hash -> cwe_id mapping from PostgreSQL COPY commands without requiring a database server."""
|
|
132
|
+
console.print("\n[bold]Parsing SQL dump metadata for CWE mappings...[/bold]")
|
|
133
|
+
cve_to_cwe: Dict[str, str] = {}
|
|
134
|
+
hash_to_cve: Dict[str, str] = {}
|
|
135
|
+
|
|
136
|
+
current_table = None
|
|
137
|
+
open_fn = gzip.open if sql_file.suffix == ".gz" else open
|
|
138
|
+
|
|
139
|
+
with open_fn(sql_file, "rt", encoding="utf-8", errors="replace") as f:
|
|
140
|
+
for line in f:
|
|
141
|
+
if line.startswith("COPY public."):
|
|
142
|
+
current_table = line.split()[1]
|
|
143
|
+
continue
|
|
144
|
+
if line.strip() == "\\.":
|
|
145
|
+
if current_table in ("public.fixes", "public.cwe_classification") and hash_to_cve and cve_to_cwe:
|
|
146
|
+
# Both key tables parsed, stop early to save processing gigabytes of raw methods
|
|
147
|
+
break
|
|
148
|
+
current_table = None
|
|
149
|
+
continue
|
|
150
|
+
|
|
151
|
+
if not current_table:
|
|
152
|
+
continue
|
|
153
|
+
|
|
154
|
+
# 1. Parse cwe_classification: (cve_id, cwe_id)
|
|
155
|
+
if current_table == "public.cwe_classification":
|
|
156
|
+
parts = line.rstrip("\n").split("\t")
|
|
157
|
+
if len(parts) >= 2:
|
|
158
|
+
cve_id = parts[0].strip()
|
|
159
|
+
cwe_id = normalize_cwe_string(parts[1].strip())
|
|
160
|
+
if cwe_id:
|
|
161
|
+
cve_to_cwe[cve_id] = cwe_id
|
|
162
|
+
|
|
163
|
+
# 2. Parse fixes: (cve_id, hash, repo_url)
|
|
164
|
+
elif current_table == "public.fixes":
|
|
165
|
+
parts = line.rstrip("\n").split("\t")
|
|
166
|
+
if len(parts) >= 2:
|
|
167
|
+
cve_id = parts[0].strip()
|
|
168
|
+
commit_hash = parts[1].strip().lower()
|
|
169
|
+
if commit_hash and cve_id:
|
|
170
|
+
hash_to_cve[commit_hash] = cve_id
|
|
171
|
+
|
|
172
|
+
# Combine into hash -> CWE
|
|
173
|
+
hash_to_cwe: Dict[str, str] = {}
|
|
174
|
+
for h, cve in hash_to_cve.items():
|
|
175
|
+
if cve in cve_to_cwe:
|
|
176
|
+
hash_to_cwe[h] = cve_to_cwe[cve]
|
|
177
|
+
|
|
178
|
+
console.print(f"[bold green]Parsed {len(hash_to_cwe)} commit-to-CWE mappings[/bold green] from {len(cve_to_cwe)} unique CVE records.")
|
|
179
|
+
return hash_to_cwe
|
|
180
|
+
|
|
181
|
+
|
|
182
|
+
def extract_curated_samples_from_patches(
|
|
183
|
+
zip_path: Path,
|
|
184
|
+
hash_to_cwe: Dict[str, str],
|
|
185
|
+
limit_per_category: int = 1500,
|
|
186
|
+
) -> List[Tuple[str, List[str]]]:
|
|
187
|
+
"""Streams unified git diffs from the patch ZIP and extracts vulnerable and benign AST slices."""
|
|
188
|
+
console.print("\n[bold]Extracting code AST slices and benign controls from patch archive...[/bold]")
|
|
189
|
+
samples: List[Tuple[str, List[str]]] = []
|
|
190
|
+
category_counts: Dict[str, int] = {cat: 0 for cat in CWE_TO_CATEGORY.values()}
|
|
191
|
+
benign_count = 0
|
|
192
|
+
max_benign = limit_per_category * 4
|
|
193
|
+
|
|
194
|
+
with zipfile.ZipFile(zip_path, "r") as zf:
|
|
195
|
+
namelist = zf.namelist()
|
|
196
|
+
console.print(f"Total patch files in archive: [bold]{len(namelist)}[/bold]")
|
|
197
|
+
|
|
198
|
+
for name in namelist:
|
|
199
|
+
if not (name.endswith(".patch") or name.endswith(".diff")):
|
|
200
|
+
continue
|
|
201
|
+
|
|
202
|
+
# Extract commit hash from filename: e.g., github.com_owner_repo_<commit_hash>.patch
|
|
203
|
+
m = re.search(r"_([0-9a-fA-F]{32,40})\.(patch|diff)", name)
|
|
204
|
+
commit_hash = m.group(1).lower() if m else ""
|
|
205
|
+
cwe_id = hash_to_cwe.get(commit_hash, "")
|
|
206
|
+
|
|
207
|
+
# If no direct hash match, check commit msg / patch content for CWE tags
|
|
208
|
+
diff_text = zf.read(name).decode("utf-8", errors="replace")
|
|
209
|
+
if not cwe_id:
|
|
210
|
+
cwe_match = re.search(r"CWE-\d+", diff_text, flags=re.IGNORECASE)
|
|
211
|
+
if cwe_match:
|
|
212
|
+
cwe_id = cwe_match.group(0).upper()
|
|
213
|
+
|
|
214
|
+
category = CWE_TO_CATEGORY.get(cwe_id)
|
|
215
|
+
if not category and ("sql" in diff_text.lower() or "exec(" in diff_text):
|
|
216
|
+
category = "INJECTION"
|
|
217
|
+
elif not category and ("tenant" in diff_text.lower() or "user_id" in diff_text.lower()):
|
|
218
|
+
category = "BOLA"
|
|
219
|
+
|
|
220
|
+
if not category:
|
|
221
|
+
category = "INJECTION"
|
|
222
|
+
|
|
223
|
+
if category_counts.get(category, 0) >= limit_per_category and benign_count >= max_benign:
|
|
224
|
+
continue
|
|
225
|
+
|
|
226
|
+
# Parse unified diff into deleted (vulnerable) and added (fixed) code lines
|
|
227
|
+
vuln_lines = []
|
|
228
|
+
fixed_lines = []
|
|
229
|
+
for line in diff_text.splitlines():
|
|
230
|
+
if line.startswith("-") and not line.startswith("---"):
|
|
231
|
+
vuln_lines.append(line[1:])
|
|
232
|
+
elif line.startswith("+") and not line.startswith("+++"):
|
|
233
|
+
fixed_lines.append(line[1:])
|
|
234
|
+
|
|
235
|
+
vuln_code = "\n".join(vuln_lines).strip()
|
|
236
|
+
fixed_code = "\n".join(fixed_lines).strip()
|
|
237
|
+
|
|
238
|
+
# Ingest vulnerable sample if meaningful length
|
|
239
|
+
if len(vuln_code) > 40 and category_counts.get(category, 0) < limit_per_category:
|
|
240
|
+
norm_vuln = normalize_code_slice(vuln_code[:1200])
|
|
241
|
+
samples.append((norm_vuln, [category]))
|
|
242
|
+
category_counts[category] = category_counts.get(category, 0) + 1
|
|
243
|
+
|
|
244
|
+
# Ingest benign fixed sample
|
|
245
|
+
if len(fixed_code) > 40 and benign_count < max_benign:
|
|
246
|
+
norm_fixed = normalize_code_slice(fixed_code[:1200])
|
|
247
|
+
samples.append((norm_fixed, []))
|
|
248
|
+
benign_count += 1
|
|
249
|
+
|
|
250
|
+
console.print(f"[bold green]Successfully extracted {len(samples)} curated code samples[/bold green] ({benign_count} benign controls):")
|
|
251
|
+
for cat, cnt in category_counts.items():
|
|
252
|
+
if cnt > 0:
|
|
253
|
+
console.print(f" • {cat:18}: [cyan]{cnt}[/cyan] samples")
|
|
254
|
+
|
|
255
|
+
return samples
|
|
256
|
+
|
|
257
|
+
|
|
258
|
+
def train_securebert_on_morefixes(
|
|
259
|
+
samples: List[Tuple[str, List[str]]],
|
|
260
|
+
output_dir: str = ".trace/models/securebert-finetuned",
|
|
261
|
+
epochs: int = 5,
|
|
262
|
+
batch_size: int = 32,
|
|
263
|
+
) -> Dict[str, float]:
|
|
264
|
+
"""Executes GPU-accelerated supervised fine-tuning of SecureBERT on RTX 4050."""
|
|
265
|
+
from trace_engine.intelligence.training.trainer import SecureBERTTrainer, TrainingConfig
|
|
266
|
+
from trace_engine.intelligence.training.dataset import VulnerabilityDataset
|
|
267
|
+
from transformers import AutoTokenizer, AutoModelForSequenceClassification
|
|
268
|
+
import torch
|
|
269
|
+
|
|
270
|
+
console.print("\n[bold green]Initializing Supervised GPU Training Engine on RTX 4050...[/bold green]")
|
|
271
|
+
config = TrainingConfig(
|
|
272
|
+
model_name="ehsanaghaei/SecureBERT",
|
|
273
|
+
epochs=epochs,
|
|
274
|
+
batch_size=batch_size,
|
|
275
|
+
learning_rate=3e-5,
|
|
276
|
+
fp16=True,
|
|
277
|
+
output_dir=output_dir,
|
|
278
|
+
)
|
|
279
|
+
trainer = SecureBERTTrainer(config)
|
|
280
|
+
|
|
281
|
+
# Re-use trainer with injected custom MoreFixes samples
|
|
282
|
+
tokenizer = AutoTokenizer.from_pretrained(config.model_name)
|
|
283
|
+
model = AutoModelForSequenceClassification.from_pretrained(
|
|
284
|
+
config.model_name,
|
|
285
|
+
num_labels=10,
|
|
286
|
+
ignore_mismatched_sizes=True,
|
|
287
|
+
).to(trainer.device)
|
|
288
|
+
|
|
289
|
+
dataset = VulnerabilityDataset(samples, tokenizer, max_length=config.max_length)
|
|
290
|
+
val_size = int(len(dataset) * 0.15)
|
|
291
|
+
train_size = len(dataset) - val_size
|
|
292
|
+
|
|
293
|
+
train_data, val_data = torch.utils.data.random_split(
|
|
294
|
+
dataset,
|
|
295
|
+
[train_size, val_size],
|
|
296
|
+
generator=torch.Generator().manual_seed(42),
|
|
297
|
+
)
|
|
298
|
+
|
|
299
|
+
train_loader = torch.utils.data.DataLoader(train_data, batch_size=config.batch_size, shuffle=True)
|
|
300
|
+
val_loader = torch.utils.data.DataLoader(val_data, batch_size=config.batch_size, shuffle=False)
|
|
301
|
+
|
|
302
|
+
pos_weights = dataset.calculate_pos_weights().to(trainer.device)
|
|
303
|
+
criterion = torch.nn.BCEWithLogitsLoss(pos_weight=pos_weights)
|
|
304
|
+
optimizer = torch.optim.AdamW(model.parameters(), lr=config.learning_rate, weight_decay=config.weight_decay)
|
|
305
|
+
scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=epochs, eta_min=5e-6)
|
|
306
|
+
|
|
307
|
+
console.print(f"\n[bold green]Fine-Tuning on {len(dataset)} Total Samples ({train_size} train, {val_size} val)[/bold green]:")
|
|
308
|
+
|
|
309
|
+
history = []
|
|
310
|
+
best_f1 = 0.0
|
|
311
|
+
|
|
312
|
+
for epoch in range(1, epochs + 1):
|
|
313
|
+
t0 = time.perf_counter()
|
|
314
|
+
model.train()
|
|
315
|
+
train_loss = 0.0
|
|
316
|
+
|
|
317
|
+
for batch in train_loader:
|
|
318
|
+
optimizer.zero_grad()
|
|
319
|
+
input_ids = batch["input_ids"].to(trainer.device)
|
|
320
|
+
mask = batch["attention_mask"].to(trainer.device)
|
|
321
|
+
labels = batch["labels"].to(trainer.device)
|
|
322
|
+
|
|
323
|
+
if trainer.scaler:
|
|
324
|
+
with torch.amp.autocast("cuda"):
|
|
325
|
+
outputs = model(input_ids=input_ids, attention_mask=mask)
|
|
326
|
+
loss = criterion(outputs.logits, labels)
|
|
327
|
+
trainer.scaler.scale(loss).backward()
|
|
328
|
+
trainer.scaler.step(optimizer)
|
|
329
|
+
trainer.scaler.update()
|
|
330
|
+
else:
|
|
331
|
+
outputs = model(input_ids=input_ids, attention_mask=mask)
|
|
332
|
+
loss = criterion(outputs.logits, labels)
|
|
333
|
+
loss.backward()
|
|
334
|
+
optimizer.step()
|
|
335
|
+
|
|
336
|
+
train_loss += loss.item()
|
|
337
|
+
|
|
338
|
+
scheduler.step()
|
|
339
|
+
avg_train_loss = round(train_loss / len(train_loader), 4)
|
|
340
|
+
|
|
341
|
+
# Validation
|
|
342
|
+
model.eval()
|
|
343
|
+
val_loss = 0.0
|
|
344
|
+
all_probs, all_labels = [], []
|
|
345
|
+
|
|
346
|
+
with torch.no_grad():
|
|
347
|
+
for v_batch in val_loader:
|
|
348
|
+
v_ids = v_batch["input_ids"].to(trainer.device)
|
|
349
|
+
v_mask = v_batch["attention_mask"].to(trainer.device)
|
|
350
|
+
v_lbl = v_batch["labels"].to(trainer.device)
|
|
351
|
+
|
|
352
|
+
if trainer.scaler:
|
|
353
|
+
with torch.amp.autocast("cuda"):
|
|
354
|
+
out = model(input_ids=v_ids, attention_mask=v_mask)
|
|
355
|
+
l = criterion(out.logits, v_lbl)
|
|
356
|
+
else:
|
|
357
|
+
out = model(input_ids=v_ids, attention_mask=v_mask)
|
|
358
|
+
l = criterion(out.logits, v_lbl)
|
|
359
|
+
|
|
360
|
+
val_loss += l.item()
|
|
361
|
+
all_probs.append(torch.sigmoid(out.logits).cpu())
|
|
362
|
+
all_labels.append(v_lbl.cpu())
|
|
363
|
+
|
|
364
|
+
avg_val_loss = round(val_loss / len(val_loader), 4)
|
|
365
|
+
cat_probs = torch.cat(all_probs, dim=0)
|
|
366
|
+
cat_labels = torch.cat(all_labels, dim=0)
|
|
367
|
+
|
|
368
|
+
preds = (cat_probs >= 0.50).float()
|
|
369
|
+
hamming_acc = round((preds == cat_labels).float().mean().item(), 4)
|
|
370
|
+
subset_acc = round((preds == cat_labels).all(dim=-1).float().mean().item(), 4)
|
|
371
|
+
|
|
372
|
+
tp = ((preds == 1) & (cat_labels == 1)).sum().item()
|
|
373
|
+
fp = ((preds == 1) & (cat_labels == 0)).sum().item()
|
|
374
|
+
fn = ((preds == 0) & (cat_labels == 1)).sum().item()
|
|
375
|
+
|
|
376
|
+
p = round(tp / (tp + fp), 3) if (tp + fp) > 0 else 0.0
|
|
377
|
+
r = round(tp / (tp + fn), 3) if (tp + fn) > 0 else 0.0
|
|
378
|
+
micro_f1 = round((2 * p * r / (p + r)), 3) if (p + r) > 0 else 0.0
|
|
379
|
+
|
|
380
|
+
dur = round(time.perf_counter() - t0, 1)
|
|
381
|
+
|
|
382
|
+
console.print(
|
|
383
|
+
f" Epoch {epoch}/{epochs}: "
|
|
384
|
+
f"Train Loss: [cyan]{avg_train_loss}[/cyan] | "
|
|
385
|
+
f"Val Loss: [yellow]{avg_val_loss}[/yellow] | "
|
|
386
|
+
f"Hamming Acc: [bold white]{round(hamming_acc * 100, 1)}%[/bold white] | "
|
|
387
|
+
f"Exact Match: [bold white]{round(subset_acc * 100, 1)}%[/bold white] | "
|
|
388
|
+
f"Precision: [green]{int(p * 100)}%[/green] | "
|
|
389
|
+
f"Recall: [green]{int(r * 100)}%[/green] | "
|
|
390
|
+
f"Micro F1: [bold green]{micro_f1}[/bold green] [{dur}s]"
|
|
391
|
+
)
|
|
392
|
+
|
|
393
|
+
history.append({
|
|
394
|
+
"epoch": epoch,
|
|
395
|
+
"train_loss": avg_train_loss,
|
|
396
|
+
"val_loss": avg_val_loss,
|
|
397
|
+
"hamming_acc": hamming_acc,
|
|
398
|
+
"exact_match_acc": subset_acc,
|
|
399
|
+
"f1": micro_f1,
|
|
400
|
+
})
|
|
401
|
+
|
|
402
|
+
if micro_f1 > best_f1:
|
|
403
|
+
best_f1 = micro_f1
|
|
404
|
+
out_p = Path(output_dir)
|
|
405
|
+
out_p.mkdir(parents=True, exist_ok=True)
|
|
406
|
+
model.save_pretrained(out_p)
|
|
407
|
+
tokenizer.save_pretrained(out_p)
|
|
408
|
+
|
|
409
|
+
console.print(f"\n[bold green]Training Complete! Best Micro-F1: {best_f1:.3f}. Checkpoint saved to {output_dir}[/bold green]")
|
|
410
|
+
return history[-1] if history else {}
|
|
411
|
+
|
|
412
|
+
|
|
413
|
+
def cleanup_dataset_directory(dataset_dir: Path) -> None:
|
|
414
|
+
"""Deletes downloaded dataset files from disk as requested by user."""
|
|
415
|
+
console.print(f"\n[bold red]Deleting raw dataset files from {dataset_dir} as requested...[/bold red]")
|
|
416
|
+
try:
|
|
417
|
+
if dataset_dir.exists():
|
|
418
|
+
for item in dataset_dir.iterdir():
|
|
419
|
+
if item.is_file():
|
|
420
|
+
item.unlink()
|
|
421
|
+
elif item.is_dir():
|
|
422
|
+
shutil.rmtree(item, ignore_errors=True)
|
|
423
|
+
console.print(f"[green]Successfully deleted raw dataset files in {dataset_dir}.[/green]")
|
|
424
|
+
except Exception as e:
|
|
425
|
+
console.print(f"[yellow]Warning during cleanup: {e}[/yellow]")
|
|
426
|
+
|
|
427
|
+
|
|
428
|
+
def run_pipeline(
|
|
429
|
+
dataset_dir: Path = Path("C:/dataset"),
|
|
430
|
+
workers: int = 16,
|
|
431
|
+
limit_per_category: int = 1500,
|
|
432
|
+
epochs: int = 5,
|
|
433
|
+
auto_delete: bool = True,
|
|
434
|
+
) -> None:
|
|
435
|
+
"""Orchestrates end-to-end download, extraction, training, and cleanup."""
|
|
436
|
+
t_start = time.perf_counter()
|
|
437
|
+
downloader = ZenodoChunkDownloader(dataset_dir, workers=workers, chunk_size_mb=8)
|
|
438
|
+
|
|
439
|
+
# 1. Download both files
|
|
440
|
+
console.print("\n[bold]Step 1/4: Downloading MoreFixes Archives to C:\\dataset\\[/bold]")
|
|
441
|
+
patch_file = downloader.download_file(PATCHES_URL, PATCHES_SIZE, "patch-files2026-06-20.zip")
|
|
442
|
+
sql_file = downloader.download_file(SQL_DUMP_URL, SQL_DUMP_SIZE, "dump-2026-06-20.sql.gz")
|
|
443
|
+
|
|
444
|
+
# 2. Extract mappings and code samples
|
|
445
|
+
console.print("\n[bold]Step 2/4: Extracting Code Slices & Ground-Truth Labels[/bold]")
|
|
446
|
+
hash_to_cwe = extract_cwe_mapping_from_sql_dump(sql_file)
|
|
447
|
+
samples = extract_curated_samples_from_patches(patch_file, hash_to_cwe, limit_per_category=limit_per_category)
|
|
448
|
+
|
|
449
|
+
# 3. Train SecureBERT with GPU acceleration
|
|
450
|
+
console.print("\n[bold]Step 3/4: Supervised GPU Fine-Tuning (RTX 4050 Laptop GPU)[/bold]")
|
|
451
|
+
train_securebert_on_morefixes(samples, epochs=epochs)
|
|
452
|
+
|
|
453
|
+
# 4. Cleanup dataset files
|
|
454
|
+
if auto_delete:
|
|
455
|
+
console.print("\n[bold]Step 4/4: Post-Training Dataset Cleanup[/bold]")
|
|
456
|
+
cleanup_dataset_directory(dataset_dir)
|
|
457
|
+
|
|
458
|
+
total_time = round(time.perf_counter() - t_start, 1)
|
|
459
|
+
console.print(f"\n[bold green]Entire MoreFixes Pipeline Completed in {total_time}s![/bold green]")
|
|
460
|
+
|
|
461
|
+
|
|
462
|
+
if __name__ == "__main__":
|
|
463
|
+
run_pipeline()
|