import numpy as np
import wave
import csv
import os

# --- Физические параметры трубы и среды (Rivixi-FakeWAV v4 - High Contrast Benchmark Setting) ---
E, RHO, NU = 200e9, 7850.0, 0.3
V = np.sqrt(E / (RHO * (1 - NU**2)))  # ~5291 м/с, скорость звука в стали

SR = 21362
DURATION = 30.0

# 1. Окрашенный стационарный шум: спад -5.7 дБ/декада, базовый уровень RMS -22.0 dBFS
TARGET_RMS_DBFS = -22.0
SPECTRAL_SLOPE = -5.7

# 2. Параметры четко сформированной яркой утечки (Экспериментальные параметры)
BURST_F_LO, BURST_F_HI = 2000.0, 3000.0
BURST_DURATION = 0.050        # 50 мс - длительность вспышки
BURST_PEAK_DBFS = -3.0        # Яркий пик (-3.0 dBFS против фона -22.0 dBFS -> контраст +19 dB!)
CONTINUOUS_JET_DBFS = -12.0   # Непрерывный шум струи для мощного пика коррелятора (+10 dB над фоном)

ATTEN_DB_PER_M = 0.03
BURSTS_PER_SEC = 2.0          # ~60 вспышек на 30-секундный файл (по 2-3 в каждое окно)

N_SAMPLES = int(SR * DURATION)
SEED = 42

def shaped_noise(n_samples, sr, slope_db_per_decade, rng):
    white = rng.standard_normal(n_samples)
    spectrum = np.fft.rfft(white)
    freqs = np.fft.rfftfreq(n_samples, d=1.0 / sr)
    freqs_safe = np.where(freqs < 1.0, 1.0, freqs)
    scale = freqs_safe ** (slope_db_per_decade / 20.0)
    scale[0] = 0.0
    return np.fft.irfft(spectrum * scale, n=n_samples)

def normalize_rms(x, target_dbfs):
    rms = np.sqrt(np.mean(x ** 2))
    if rms < 1e-9: return x
    target = 10 ** (target_dbfs / 20.0)
    return x * (target / rms)

def make_background(n, sr, rng):
    left = normalize_rms(shaped_noise(n, sr, SPECTRAL_SLOPE, rng), TARGET_RMS_DBFS)
    right = normalize_rms(shaped_noise(n, sr, SPECTRAL_SLOPE, rng), TARGET_RMS_DBFS)
    return left, right

def make_burst(sr, rng):
    n = int(sr * BURST_DURATION)
    white = rng.standard_normal(n)
    spec = np.fft.rfft(white)
    freqs = np.fft.rfftfreq(n, d=1.0 / sr)
    band = (freqs >= BURST_F_LO) & (freqs <= BURST_F_HI)
    spec[~band] = 0.0
    burst = np.fft.irfft(spec, n=n)
    
    # Огибающая: быстрая атака (15%) и экспоненциальный спад
    t = np.linspace(0, 1, n)
    envelope = np.where(t < 0.15, t / 0.15, np.exp(-4.0 * (t - 0.15)))
    burst *= envelope
    max_val = np.max(np.abs(burst))
    if max_val > 1e-9:
        burst /= max_val
    return burst

def make_continuous_jet_noise(n, sr, rng):
    white = rng.standard_normal(n)
    spec = np.fft.rfft(white)
    freqs = np.fft.rfftfreq(n, d=1.0 / sr)
    band = (freqs >= BURST_F_LO) & (freqs <= BURST_F_HI)
    spec[~band] = 0.0
    jet = np.fft.irfft(spec, n=n)
    return normalize_rms(jet, CONTINUOUS_JET_DBFS)

def fft_fractional_delay(x, tau, sr):
    n = len(x)
    n_pad = n * 2
    spec = np.fft.rfft(x, n=n_pad)
    freqs = np.fft.rfftfreq(n_pad, d=1.0 / sr)
    phase_shift = np.exp(-1j * 2 * np.pi * freqs * tau)
    shifted = np.fft.irfft(spec * phase_shift, n=n_pad)
    return shifted[:n]

def insert_burst(channel, burst, sample_pos, gain):
    n = len(burst)
    end = sample_pos + n
    if end > len(channel):
        n = len(channel) - sample_pos
        if n <= 0: return
        burst = burst[:n]
        end = sample_pos + n
    channel[sample_pos:end] += burst * gain

def generate_file(out_path, pipe_length, is_leak, rng):
    left, right = make_background(N_SAMPLES, SR, rng)
    defect_records = []
    
    if is_leak:
        x = rng.uniform(5.0, pipe_length - 5.0)
        d_left = x
        d_right = pipe_length - x
        
        # Точный TDOA сдвиг
        dt_left = d_left / V
        dt_right = d_right / V
        delta_t = dt_left - dt_right
        
        # 1. Непрерывный шум струи со сдвигом TDOA
        jet_raw = make_continuous_jet_noise(N_SAMPLES, SR, rng)
        jet_left = jet_raw * (10.0 ** (-ATTEN_DB_PER_M * d_left / 20.0))
        jet_right = fft_fractional_delay(jet_raw, delta_t, SR) * (10.0 ** (-ATTEN_DB_PER_M * d_right / 20.0))
        left += jet_left
        right += jet_right
        
        # 2. Серия ярких вспышек со сдвигом TDOA
        gain_left_linear = 10.0 ** (BURST_PEAK_DBFS / 20.0) * (10.0 ** (-ATTEN_DB_PER_M * d_left / 20.0))
        gain_right_linear = 10.0 ** (BURST_PEAK_DBFS / 20.0) * (10.0 ** (-ATTEN_DB_PER_M * d_right / 20.0))
        
        n_bursts = int(DURATION * BURSTS_PER_SEC)
        burst_times = np.linspace(0.2, DURATION - 0.2, n_bursts) + rng.uniform(-0.05, 0.05, n_bursts)
        
        for t_b in burst_times:
            burst_raw = make_burst(SR, rng)
            burst_left = burst_raw
            burst_right = fft_fractional_delay(burst_raw, delta_t, SR)
            
            pos_left = int(t_b * SR)
            pos_right = int(t_b * SR)
            
            if 0 <= pos_left < N_SAMPLES:
                insert_burst(left, burst_left, pos_left, gain_left_linear)
            if 0 <= pos_right < N_SAMPLES:
                insert_burst(right, burst_right, pos_right, gain_right_linear)
                
        defect_records.append({
            'defect_index': 0,
            'pos_m': round(x, 2),
            'd_left_m': round(d_left, 2),
            'd_right_m': round(d_right, 2),
            'gain_left_dbfs': round(20*np.log10(max(gain_left_linear, 1e-6)), 2),
            'gain_right_dbfs': round(20*np.log10(max(gain_right_linear, 1e-6)), 2),
            'delta_t_ms': round(delta_t * 1000.0, 4)
        })
    
    # Защита от клиппинга
    max_amp = max(np.max(np.abs(left)), np.max(np.abs(right)))
    if max_amp > 0.99:
        left = left / max_amp * 0.95
        right = right / max_amp * 0.95
        
    audio_int16 = np.zeros((N_SAMPLES, 2), dtype=np.int16)
    audio_int16[:, 0] = np.clip(left * 32767.0, -32768, 32767).astype(np.int16)
    audio_int16[:, 1] = np.clip(right * 32767.0, -32768, 32767).astype(np.int16)
    
    with wave.open(out_path, 'wb') as wf:
        wf.setnchannels(2)
        wf.setsampwidth(2)
        wf.setframerate(SR)
        wf.writeframes(audio_int16.tobytes())
        
    return defect_records

def main():
    rng = np.random.default_rng(SEED)
    
    out_leak_dir = r'C:\Users\ivaew\Desktop\RIVIXI\SAAS9\1d_cnn_degradation\synthetic\leak'
    out_norm_dir = r'C:\Users\ivaew\Desktop\RIVIXI\SAAS9\1d_cnn_degradation\synthetic\normal'
    doc_dir = r'C:\Users\ivaew\Desktop\RIVIXI\SAAS9\2d_cnn_degradation\Допиливаем статью'
    
    os.makedirs(out_leak_dir, exist_ok=True)
    os.makedirs(out_norm_dir, exist_ok=True)
    
    manifest_rows = []
    
    print("Генерация 65 файлов утечек (экспериментальные параметры: RMS -22.0 dBFS, вспышки -3.0 dBFS, струя -12.0 dBFS)...")
    for i in range(1, 66):
        fn = f"synth_leak_{i:03d}.wav"
        p = os.path.join(out_leak_dir, fn)
        pipe_l = rng.uniform(50.0, 150.0)
        recs = generate_file(p, pipe_l, is_leak=True, rng=rng)
        rec = recs[0]
        manifest_rows.append({
            'filename': fn,
            'class': 'leak',
            'pipe_length_m': round(pipe_l, 2),
            'defect_x_m': rec['pos_m'],
            'delta_t_ms': rec['delta_t_ms'],
            'gain_left_dbfs': rec['gain_left_dbfs'],
            'gain_right_dbfs': rec['gain_right_dbfs']
        })
        
    print("Генерация 30 файлов нормы (RMS -22.0 dBFS)...")
    for i in range(1, 31):
        fn = f"synth_normal_{i:03d}.wav"
        p = os.path.join(out_norm_dir, fn)
        pipe_l = rng.uniform(50.0, 150.0)
        generate_file(p, pipe_l, is_leak=False, rng=rng)
        manifest_rows.append({
            'filename': fn,
            'class': 'normal',
            'pipe_length_m': round(pipe_l, 2),
            'defect_x_m': 'N/A',
            'delta_t_ms': 'N/A',
            'gain_left_dbfs': 'N/A',
            'gain_right_dbfs': 'N/A'
        })
        
    manifest_csv = os.path.join(doc_dir, 'synthetic_manifest_experiment.csv')
    with open(manifest_csv, 'w', newline='', encoding='utf-8') as f:
        writer = csv.DictWriter(f, fieldnames=['filename', 'class', 'pipe_length_m', 'defect_x_m', 'delta_t_ms', 'gain_left_dbfs', 'gain_right_dbfs'])
        writer.writeheader()
        writer.writerows(manifest_rows)
        
    print("Файлы эксперимента успешно восстановлены!")

if __name__ == '__main__':
    main()
