Repository navigation
Expand file tree
/
Copy pathcvefix_formatter.py
More file actions
290 lines (260 loc) · 11.7 KB
/
Copy pathcvefix_formatter.py
File metadata and controls
290 lines (260 loc) · 11.7 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
import os
import ast
import typer
import time
from loguru import logger
from pathlib import Path
from multiprocessing import Process, Queue
from sqlalchemy import create_engine, func
from sqlalchemy.orm import sessionmaker
from models.meta import CVEMeta, PatchMeta
from models.cvefix import CVE, CWE, Fixes, Repository, Commit, FileChange, CWEClassification
from utils.validate_path import validate_path_db
from utils.mongo_interface import MongoInterface
from models.desc import CVEDescription, SemanticDesc
from vulnsrc_formatter import ask_llm
BATCH_SIZE = 10
STOP_SIGNAL = object() # Sentinel object to stop the queue
def do_query(session: sessionmaker, rec_cnt: int = None):
"""
Do the query and return the results.
Args:
session: The SQLAlchemy session.
rec_cnt: The number of records to return. If None, return all records.
Returns:
The results of the query. Should be a generator.
"""
# Subquery to find CVE_ids with multiple commits
subquery_cve = (
session.query(Fixes.cve_id)
.group_by(Fixes.cve_id)
.having(func.count(Fixes.hash) > 1)
.subquery()
)
# Subquery to find commit hashes with multiple file changes
subquery_commit = (
session.query(FileChange.hash)
.group_by(FileChange.hash)
.having(func.count(FileChange.filename) > 1)
.subquery()
)
query = (
session.query(
CVE.cve_id.label("CVE_id"),
func.group_concat(CWE.cwe_id, ',').label('CWE_ids'),
CVE.description.label("CVE_description"),
Fixes.hash.label("patch_sha"),
Repository.repo_name.label("repo_name"),
Commit.msg.label("commit_message"),
FileChange.filename.label("filename"),
FileChange.code_before.label("code_before"),
FileChange.code_after.label("code_after"),
FileChange.diff.label("diff"),
)
.join(CWEClassification, CVE.cve_id == CWEClassification.cve_id)
.join(CWE, CWEClassification.cwe_id == CWE.cwe_id)
.join(Fixes, CVE.cve_id == Fixes.cve_id)
.join(Commit, (Fixes.hash == Commit.hash))
.join(Repository, Commit.repo_url == Repository.repo_url)
.join(FileChange, Commit.hash == FileChange.hash)
# NOTE: Drop CVEs with multiple commits
.filter(~CVE.cve_id.in_(subquery_cve.select()))
# NOTE: Drop commits with multiple file changes
.filter(~FileChange.hash.in_(subquery_commit.select()))
.group_by(CVE.cve_id, CVE.description, Fixes.hash, Repository.repo_name, Commit.msg, FileChange.filename, FileChange.code_before, FileChange.code_after, FileChange.diff)
).yield_per(BATCH_SIZE)
count = 0
for record in query:
yield record # Yield the record to save memory
count += 1
if rec_cnt is not None and count >= rec_cnt:
break
def fetch_cve_data(db_path: Path, query_queue: Queue, batch_cnt: int):
"""
分批次从数据库中提取信息并检查,丢弃不符合要求的记录
Args:
db_path: 数据库文件路径
Returns:
ClassCVEMeta to MQ"""
engine = create_engine(f'sqlite:///{db_path}')
Session = sessionmaker(bind=engine)
session = Session()
try:
# Call the query function
for rec in do_query(session, batch_cnt*BATCH_SIZE if batch_cnt != -1 else None):
# For each record, process it.
cve_meta = process_row(rec)
if cve_meta is not None:
logger.info(f"Put {cve_meta.cve_number} to the queue")
query_queue.put(cve_meta)
else:
logger.info("Dropped the invaild record")
finally:
# Use a dedicated sentinel object for better clarity
session.close()
query_queue.put(STOP_SIGNAL)
logger.info("Sent stop signal")
def process_row(row) -> CVEMeta:
"""
Process a row from the database and return a CVEMeta object. If the row in invaild or has data's value is None, drop it and return None.
Before instantiating the CVEMeta object, we need to finish the following steps:
1. Check if the record already exists in the MongoDB[cvefix] and [vulnsrc].
2. Check if the record has missing fields.
3. Check if the changed files have been stored in the GridFS.
Args:
row: The row from the database.
Returns:
The CVEMeta object.
"""
if row is None:
logger.warning("Empty row, skipping")
return None
# Check if the record already exists in the MongoDB by CVE_id
try:
if mongo_client.find_one({"cve_meta.cve_number": row.CVE_id}):
logger.info(
f"Record {row.CVE_id} already exists in the MongoDB cvefix. Skipping")
return None
if mongo_client_vulnsrc.find_one({"cve_meta.cve_number": row.CVE_id}):
# Should be cve_meta.cve_number
logger.info(
f"Record {row.CVE_id} already exists in the MongoDB vulnsrc. Skipping")
return None
except Exception as e:
logger.warning(f"Error {e} while checking {
row.CVE_id} in the MongoDB. Check the record manually")
return None
# Save the files to GridFS
try:
filename = row.filename
code_before = row.code_before
code_after = row.code_after
diff = row.diff
# NOTE: We assume that the insert operation is atomic so that we will not check the existence of the file knowing that the CVE is not exist in the MongoDB, to prevent the same file was found multi CVEs.
# Parse the filename
basename, extension_name = os.path.splitext(filename)
# Save the file to GridFS
code_before_file_id = mongo_client.insert_file_by_code(
code_before, f"{basename}_before{extension_name}")
code_after_file_id = mongo_client.insert_file_by_code(
code_after, f"{basename}_after{extension_name}")
diff_file_id = mongo_client.insert_file_by_code(
diff, f"{basename}.diff")
except Exception as e:
logger.warning(f"Error {e} while saving the file.")
# NOTE: The saved files before the error will not be deleted, so they may be orphaned
return None
logger.info(f"Saved files for {row.CVE_id}")
# Instantiate the CVEMeta object
try:
cve_meta = CVEMeta(
cve_number=row.CVE_id,
# NOTE: CWE_ids is a string of comma-separated CWE IDs. Noted that a CVE may have multiple CWEs.
weaknesses=row.CWE_ids.split(','),
# NOTE: CVEFIX model has no title field. Use CVE_id as a placeholder.
title=row.CVE_id,
# NOTE: The description is a JSON object, so we need to extract the English version. If unexpected, use an empty string and the record will be dropped.
description=parse_cve_desc(row.CVE_description),
patch_meta=PatchMeta(
commit_sha=row.patch_sha,
commit_message=row.commit_message,
repo=row.repo_name,
vulnerable_codes_id=[code_before_file_id],
patched_code_id=[code_after_file_id],
diff_id=[diff_file_id]
)
)
except Exception as e:
logger.warning(f"Error {e} while processing {row.CVE_id}")
return None
if cve_meta.has_missing_fields():
logger.info(f"Missing fields in {
cve_meta.cve_number}, dropping the record")
return None
logger.info(f"Processed {cve_meta.cve_number}")
return cve_meta
def parse_cve_desc(cve_desc: str) -> str:
"""
Convert the CVE description to a dictionary and extract the English version. If unexisted, return an empty string."""
try:
# Convert the string to a dictionary using ast.literal_eval instead of json.loads to prevent security issues
parsed_data = ast.literal_eval(cve_desc)
for item in parsed_data:
if isinstance(item, dict) and item.get('lang') == 'en':
return item.get('value', '') # 返回 'value' 字段或空字符串
except (ValueError, SyntaxError):
logger.warning("Failed to parse the CVE description")
return '' # NOTE: Return an empty string if failed
def extract_data(query_queue: Queue, mongo_host: str, mongo_port: int):
"""
Extract data from the queue and write it to a file.
Args:
query_queue: The queue from which to extract data.
mongo_host: The MongoDB host address.
mongo_port: The MongoDB port number.
"""
logger.info("Ready to extracting data from the queue")
valid_records_count = 0
while True:
data = query_queue.get() # Block until data is available
if type(data) is not CVEMeta:
logger.info(f"Received stop signal {type(data)}")
break # Stop the loop
cve_meta = data
# Wrap the metadata with CVEDesc class
cve_desc = CVEDescription(cve_meta=cve_meta, desc=SemanticDesc())
cve_desc = ask_llm(cve_desc, mongo_client)
if cve_desc is None:
logger.warning(
f"Failed to extract semantic information for {cve_meta.cve_number}.")
continue
# Save the data to the MongoDB
cve_desc_ready_to_insert = cve_desc.model_dump()
try:
mongo_client.insert(cve_desc_ready_to_insert)
logger.info(
f"Inserted {cve_desc.cve_meta.cve_number} to MongoDB Successfully")
valid_records_count += 1
except Exception as e:
logger.exception(f"Failed to insert {
cve_desc.cve_meta.cve_number}: {e}")
logger.info(f"The process of extracting data is finished. Total valid records processed: {
valid_records_count}")
return valid_records_count
def validate_batch_cnt(value: int) -> int:
"""
Validate the batch count value.
Args:
value: The batch count value.
Returns:
The validated batch count value.
"""
if value < -1 or value == 0:
raise typer.BadParameter(
"Batch count must be greater than or equal to -1 and cannot be 0")
return value
def main(db_path: str = typer.Argument(..., callback=validate_path_db, help="Path to the CVEfixes database file(.db)"), batch_cnt: int = typer.Option(-1, help="Number of batches to process. Default is -1, which means all batches will be processed.", callback=validate_batch_cnt), mongo_host: str = typer.Option("localhost", help="MongoDB host address"), mongo_port: int = typer.Option(27017, help="MongoDB port number")):
logger.info(f"Fetching data from CVEFix Dataset {db_path}")
# FIXME: UserWarning: MongoClient opened before fork. May not be entirely fork-safe, proceed with caution. See PyMongo's documentation for details: https://www.mongodb.com/docs/languages/python/pymongo-driver/current/faq/#is-pymongo-fork-safe-
global mongo_client, mongo_client_vulnsrc
start_time = time.time()
query_queue = Queue()
select_process = Process(target=fetch_cve_data,
args=(db_path, query_queue, batch_cnt))
extract_process = Process(target=extract_data,
args=(query_queue, mongo_host, mongo_port))
# Connect to the mongo database
mongo_client = MongoInterface(mongo_host, mongo_port, "cvefix", "cve")
mongo_client_vulnsrc = MongoInterface(
mongo_host, mongo_port, "vulnsrc", "cves")
# Start both processes
select_process.start()
extract_process.start()
# Join both processes
select_process.join()
extract_process.join()
end_time = time.time()
logger.info(f"All data fetched from CVEFix Dataset {
db_path}. Total time: {int(end_time - start_time)} seconds.")
if __name__ == "__main__":
typer.run(main)