Files
tianxuan/scripts/gen_test_data.py

293 lines
14 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""Generate synthetic TLS flow test data matching TlsDB.csv column spec.
Outputs 1998 CSV files (0.csv .. 1997.csv) with realistic TLS flow data.
Columns match TlsDB.csv exactly; per-file column order is randomized.
Columns that are all-blank in a given file are auto-dropped from that file's header.
Usage:
runtime\\python\\python.exe scripts\\gen_test_data.py [--files N] [--rows N] [--output-dir DIR]
"""
import random
import csv
import argparse
from pathlib import Path
from datetime import datetime, timedelta
random.seed(42)
# ── Column spec (exactly matches TlsDB.csv — 'row' removed because not in TlsDB) ──
COLUMNS = [
':ips', ':ipd', ':prs', ':prd', 'scnt', 'dcnt',
':ips.latd', ':ips.lond', ':ipd.latd', ':ipd.lond',
':ips.ispn', ':ipd.ispn', ':ips.orgn', ':ipd.orgn',
':ips.city', ':ipd.city', ':ips.anon', ':ipd.anon',
':ips.doma', ':ipd.doma',
'server-ip', 'client-ip', 'time', 'timestamp',
'1ipp', '4dbn', 'tabl', '4ksz', 'cnrs', 'isrs',
'cnam', '0ver', 'snam', '4dur', '8seq', '2tmo', 'name',
'source-node', 'cipher-suite', 'ecdhe-named-curve',
'0cph', '0crv', '0rnd', '0rnt',
'8ack', '8pak', '8did', '4srs', '8ppk', '8ses', '8byt', '8dbd',
'crcc', 'orga', 'orgu', 'eiph', '@iph',
]
# ── Data pools ───────────────────────────────────────────────────────────
TLS_VERSIONS = ['03 03', '03 04', '02 00', '03 01', '03 02']
TLS_WEIGHTS = [0.55, 0.30, 0.05, 0.05, 0.05]
CIPHER_HEX = ['c0 2b', 'c0 2f', '13 01', '13 02', 'c0 2c', 'cc a9', 'c0 23',
'c0 27', 'c0 13', 'c0 14', '00 9e', '00 9f', '00 35']
CIPHER_NAMES = [
'TLS_ECDHE_ECDSA_AES128_GCM_SHA256', 'TLS_ECDHE_RSA_AES128_GCM_SHA256',
'TLS_AES_128_GCM_SHA256', 'TLS_AES_256_GCM_SHA384',
'TLS_ECDHE_ECDSA_AES256_GCM_SHA384', 'TLS_ECDHE_ECDSA_CHACHA20_POLY1305',
'TLS_ECDHE_RSA_AES256_GCM_SHA384', 'TLS_ECDHE_RSA_AES128_SHA256',
'TLS_ECDHE_ECDSA_AES128_SHA', 'TLS_ECDHE_RSA_AES128_SHA',
'TLS_RSA_WITH_AES_128_GCM_SHA256', 'TLS_RSA_WITH_AES_256_GCM_SHA384',
'TLS_RSA_WITH_AES_128_CBC_SHA',
]
NAMED_CURVES_HEX = ['00 1d', '00 17', '00 18', '00 19', '00 1e']
NAMED_CURVES_TEXT = ['secp256r1', 'secp384r1', 'secp521r1', 'x25519', 'x448']
SRC_IPS = [f'{random.randint(1,223)}.{random.randint(0,255)}.{random.randint(0,255)}.{random.randint(1,254)}' for _ in range(200)]
DST_IPS = [f'{random.randint(1,223)}.{random.randint(0,255)}.{random.randint(0,255)}.{random.randint(1,254)}' for _ in range(500)]
PORTS = [443, 80, 8080, 8443, 465, 993, 995, 53, 22, 3389, 25, 110, 143, 21, 12345, 50000, 443, 443, 443, 443]
COUNTRIES = ['US','CN','KR','JP','GB','DE','FR','RU','BR','IN','SG','NL','CA','AU','HK','UA','IL','SE','NO','FI']
ISPS = ['China Telecom','China Mobile','BT Group','Orange','Deutsche Telekom','AT&T','Verizon','Comcast','NTT','KDDI','SK Telecom','Singtel','Telstra']
ORGS = ['Baidu Inc.','Alibaba Inc.','Amazon.com Inc.','Google LLC','Microsoft Corp.','Meta Platforms','Tencent','Samsung','Apple Inc.','Netflix Inc.']
CITIES = ['Beijing','Shanghai','Seoul','Berlin','Paris','London','New York','San Francisco','Tokyo','Singapore','Sydney','Moscow','Seattle','Dublin','Mumbai']
SERVICES = ['mail.example.com','api.example.com','cdn.example.com','auth.example.net','stream.example.org','login.live.com','*.cloudfront.net','*.s3.amazonaws.com','graph.facebook.com','www.google.com']
SRC_NODES = ['packet_capture','ssl_logs','netflow','zeek','suricata']
NODE_NAMES = ['wan-link','core-02','gw-09','edge-01','backbone-03']
def load_coverage(tlsdb_path=None):
"""Parse TlsDB.csv and return {column_name: coverage_float}.
Locates TlsDB.csv relative to the scripts/ directory (project root).
"""
if tlsdb_path is None:
tlsdb_path = Path(__file__).parent.parent / 'TlsDB.csv'
coverage = {}
with open(tlsdb_path, 'r', encoding='utf-8') as f:
reader = csv.DictReader(f)
for row in reader:
name = row['标题'].strip()
pct_str = row.get('覆盖率', '100%').strip().rstrip('%')
try:
coverage[name] = float(pct_str) / 100.0
except ValueError:
coverage[name] = 1.0
return coverage
COVERAGE = load_coverage()
def _rand_ip(pool):
return random.choice(pool)
def _rand_hex(n):
return ' '.join(f'{random.randint(0,255):02x}' for _ in range(n))
def _maybe(rate):
"""Return True with given probability (0.0-1.0)."""
return random.random() < rate
def _blank_or(rate, value, blank_marker=''):
"""Return value with rate probability, else blank_marker."""
return value if _maybe(rate) else blank_marker
def _blank_or_plus(rate, value):
"""Return value with rate probability, else '+' (null marker)."""
return value if _maybe(rate) else '+'
def generate_row(row_id, base_time):
"""Generate a single TLS flow row matching all TlsDB.csv columns.
row_id is a global sequential counter used to derive:
- time: base_time.strftime('%Y-%m-%d %H:%M:%S.') + f'{row_id % 1000000:06d}'
- timestamp: round(base_time.timestamp() + row_id * 0.000001, 6)
"""
src_ip = _rand_ip(SRC_IPS)
dst_ip = _rand_ip(DST_IPS)
src_port = random.choice(PORTS)
dst_port = random.choice(PORTS)
src_country = random.choice(COUNTRIES)
dst_country = random.choice(COUNTRIES)
server_ip = _rand_ip(DST_IPS + SRC_IPS)
client_ip = src_ip if _maybe(0.85) else _rand_ip(SRC_IPS)
tls_ver = random.choices(TLS_VERSIONS, weights=TLS_WEIGHTS, k=1)[0]
cipher_hex = random.choice(CIPHER_HEX)
cipher_name = random.choice(CIPHER_NAMES)
curve_hex = random.choice(NAMED_CURVES_HEX)
curve_name = random.choice(NAMED_CURVES_TEXT)
service = random.choice(SERVICES)
cert_cn = service if _maybe(0.7) else f'*.{random.choice(SERVICES).split(".")[-2:]}'
# (a) time: YYYY-MM-DD HH:MM:SS.XXXXXX
time_str = base_time.strftime('%Y-%m-%d %H:%M:%S.') + f'{row_id % 1000000:06d}'
# (b) timestamp: float with 6 decimals
ts = round(base_time.timestamp() + row_id * 0.000001, 6)
# Row-local lat/lon with realistic blank rates
src_lat = round(random.uniform(-60, 60), 6)
src_lon = round(random.uniform(-180, 180), 6)
dst_lat = round(random.uniform(-60, 60), 6)
dst_lon = round(random.uniform(-180, 180), 6)
return {
# ── Core (always present) ──
':ips': src_ip,
':ipd': dst_ip,
':prs': src_port,
':prd': dst_port,
'time': time_str,
'timestamp': ts,
# ── TLS fingerprint (very high coverage) ──
'0ver': tls_ver,
'0cph': cipher_hex,
'cipher-suite': _blank_or(COVERAGE.get('cipher-suite', 1.0), cipher_name),
'ecdhe-named-curve': _blank_or(COVERAGE.get('ecdhe-named-curve', 1.0), curve_name),
'0crv': _blank_or(COVERAGE.get('0crv', 1.0), curve_hex),
'0rnd': _blank_or(COVERAGE.get('0rnd', 1.0), _rand_hex(28)),
'0rnt': _blank_or(COVERAGE.get('0rnt', 1.0), _rand_hex(4)),
'4ksz': _blank_or(COVERAGE.get('4ksz', 1.0), str(random.choice([128, 256, 384, 512]))),
# ── Session / connection info (medium-high coverage) ──
'snam': _blank_or(COVERAGE.get('snam', 1.0), service),
'cnam': _blank_or(COVERAGE.get('cnam', 1.0), cert_cn),
'server-ip': _blank_or(COVERAGE.get('server-ip', 1.0), server_ip),
'client-ip': _blank_or(COVERAGE.get('client-ip', 1.0), client_ip),
'4dur': _blank_or(COVERAGE.get('4dur', 1.0), str(round(random.uniform(0.5, 120.0), 2))),
'8ack': _blank_or(COVERAGE.get('8ack', 1.0), str(random.randint(100, 200000))),
'8pak': _blank_or(COVERAGE.get('8pak', 1.0), str(random.randint(1, 2000))),
'8ppk': _blank_or(COVERAGE.get('8ppk', 1.0), str(random.randint(1, 500))),
'8seq': _blank_or(COVERAGE.get('8seq', 1.0), str(round(random.uniform(1, 2000), 4))),
'8ses': _blank_or(COVERAGE.get('8ses', 1.0), str(round(random.uniform(1, 2000), 4))),
'8byt': _blank_or(COVERAGE.get('8byt', 1.0), str(random.randint(200, 500000))),
'2tmo': _blank_or(COVERAGE.get('2tmo', 1.0), str(round(random.uniform(0, 1000), 4))),
# ── TCP / network info (medium coverage) ──
'1ipp': str(random.randint(1, 50)),
'tabl': random.choice(['TlsC', 'TlsS']),
'name': _blank_or(COVERAGE.get('name', 1.0), random.choice(NODE_NAMES)),
'source-node': _blank_or(COVERAGE.get('source-node', 1.0), random.choice(SRC_NODES)),
# ── Session resumption (boolean) ──
'cnrs': _blank_or(COVERAGE.get('cnrs', 1.0), '+'),
'isrs': _blank_or(COVERAGE.get('isrs', 1.0), '+'),
# ── Database metadata (low-medium coverage) ──
'4dbn': _blank_or(COVERAGE.get('4dbn', 1.0), str(random.randint(1, 100))),
'8dbd': _blank_or(COVERAGE.get('8dbd', 1.0), str(random.randint(1, 200))),
'8did': _blank_or(COVERAGE.get('8did', 1.0), str(random.randint(1000, 9999))),
'4srs': _blank_or(COVERAGE.get('4srs', 1.0), str(random.randint(10000, 99999))),
# ── GeoIP location ──
':ips.latd': _blank_or_plus(COVERAGE.get(':ips.latd', 1.0), str(src_lat)),
':ips.lond': _blank_or_plus(COVERAGE.get(':ips.lond', 1.0), str(src_lon)),
':ipd.latd': _blank_or_plus(COVERAGE.get(':ipd.latd', 1.0), str(dst_lat)),
':ipd.lond': _blank_or_plus(COVERAGE.get(':ipd.lond', 1.0), str(dst_lon)),
':ips.ispn': _blank_or(COVERAGE.get(':ips.ispn', 1.0), random.choice(ISPS)),
':ipd.ispn': _blank_or(COVERAGE.get(':ipd.ispn', 1.0), random.choice(ISPS)),
':ips.orgn': _blank_or(COVERAGE.get(':ips.orgn', 1.0), random.choice(ORGS)),
':ipd.orgn': _blank_or(COVERAGE.get(':ipd.orgn', 1.0), random.choice(ORGS)),
':ips.city': _blank_or(COVERAGE.get(':ips.city', 1.0), random.choice(CITIES)),
':ipd.city': _blank_or(COVERAGE.get(':ipd.city', 1.0), random.choice(CITIES)),
':ips.anon': _blank_or(COVERAGE.get(':ips.anon', 1.0), 'anon, hosting'),
':ipd.anon': _blank_or(COVERAGE.get(':ipd.anon', 1.0), 'anon, hosting'),
':ips.doma': _blank_or(COVERAGE.get(':ips.doma', 1.0), random.choice(SERVICES)),
':ipd.doma': _blank_or(COVERAGE.get(':ipd.doma', 1.0), random.choice(SERVICES)),
# ── Two-letter / abbreviation columns ──
# (c) scnt/dcnt: .lower() → e.g. 'us', 'cn'
'scnt': _blank_or(COVERAGE.get('scnt', 1.0), src_country.lower()),
'dcnt': _blank_or(COVERAGE.get('dcnt', 1.0), dst_country.lower()),
'crcc': _blank_or(COVERAGE.get('crcc', 1.0), random.choice(['OK', 'ER', '--'])),
'@iph': _blank_or(COVERAGE.get('@iph', 1.0), random.choice(['A', 'B', 'C', 'D'])),
'eiph': _blank_or_plus(COVERAGE.get('eiph', 1.0), '+'),
# ── Obscure / rarely populated fields ──
'orga': _blank_or(COVERAGE.get('orga', 1.0), random.choice(ORGS)),
'orgu': _blank_or(COVERAGE.get('orgu', 1.0), str(random.randint(1000, 99999))),
}
def main():
parser = argparse.ArgumentParser(
description='Generate TLS flow test data matching TlsDB.csv — '
'1998 files (0.csv..1997.csv) with randomized column order per file'
)
parser.add_argument('--files', type=int, default=1998,
help='Number of CSV files (default: 1998)')
parser.add_argument('--rows', type=int, default=100,
help='Rows per file (default: 100)')
parser.add_argument('--output-dir', type=str, default='data/test_csvs',
help='Output directory (default: data/test_csvs)')
parser.add_argument('--seed', type=int, default=42,
help='Random seed for data generation (default: 42)')
args = parser.parse_args()
random.seed(args.seed)
out_dir = Path(args.output_dir)
out_dir.mkdir(parents=True, exist_ok=True)
base_time = datetime(2026, 6, 1, 0, 0, 0)
total_rows = 0
for file_idx in range(args.files):
# File-specific RNG for column-order randomization (g)
file_seed = args.seed + file_idx
file_rng = random.Random(file_seed)
# Generate all rows for this file
rows = []
for row_i in range(args.rows):
row = generate_row(total_rows + row_i, base_time)
rows.append(row)
# (f) Auto-drop columns that are ALL blank ('' or '+') across every row
non_blank_cols = set()
for row in rows:
for col in COLUMNS:
val = row.get(col, '')
if val != '' and val != '+':
non_blank_cols.add(col)
file_columns = [c for c in COLUMNS if c in non_blank_cols]
# (g) Randomize column order per file — high+low coverage interleaved
shuffled = file_columns[:]
file_rng.shuffle(shuffled)
# (e) Write: 0.csv, 1.csv, ... 1997.csv
out_path = out_dir / f'{file_idx}.csv'
with open(out_path, 'w', newline='', encoding='utf-8') as f:
writer = csv.DictWriter(f, fieldnames=shuffled, extrasaction='ignore')
writer.writeheader()
for row in rows:
writer.writerow(row)
total_rows += args.rows
if (file_idx + 1) % 100 == 0 or file_idx == 0 or file_idx == args.files - 1:
print(f' [{file_idx + 1}/{args.files}] {total_rows} rows, '
f'{len(shuffled)} cols, seed={file_seed}')
print(f'\nDone: {args.files} files × {args.rows} rows = {total_rows} total rows')
print(f'Output: {out_dir.resolve()}')
if __name__ == '__main__':
main()