# samtools view -h input_namesorted.bam | python3 filter-BAM-SOI-all-alignments.py /mounts/lovelace/temporary/SoI-for-workflow.csv | samtools view -b -o filtered_taxids_complete.bam

import sys
import csv
from functools import lru_cache
import psycopg2

# load the species of interest CSV (derived from the Drive Sheet), taxid is in the 4th column 
target_taxids = set()
csv_file_path = sys.argv[1]

with open(csv_file_path, mode="r", encoding="utf-8") as f:
  reader = csv.reader(f)
  for row_idx, row in enumerate(reader):
    # Automatically skips header if the fourth column value isn't a numeric taxid
    if row_idx == 0 and not row[3].strip().isdigit():
      continue
    if len(row) > 3:
      taxid = row[3].strip()
      if taxid:
        target_taxids.add(taxid)

  sys.stderr.write(
      f"[INFO] Loaded {len(target_taxids)} unique target taxids from CSV.\n"
  )
  if len(target_taxids) > 0:
    sys.stderr.write(
        f"[INFO] Sample target taxids: {list(target_taxids)[:5]}\n"
    )
  else:
    sys.stderr.write(
        "[WARNING] Target taxids set is EMPTY! Check your CSV column index (0-indexed 3 = 4th column).\n"
    )

db_connection = psycopg2.connect(dbname="ncbi", user="fieldsci", password="skalanes")

# @lru_cache(maxsize=500000)
def get_taxid_for_accession(accession, db_connection):
    try:
        with db_connection.cursor() as cur:
            cur.execute("""
                SELECT tax_id
                FROM accession_taxid
                WHERE accession_version = %s
                LIMIT 1;
            """, (accession,))
            row = cur.fetchone()
  
            if row:
                return str(row[0])
            else:
                return None

    except Exception as e:
        print(f"[Error] {e}")
        return None

current_qname = None
current_chunk = []
keep_read = False
debug_lookups = 0

def process_chunk(chunk, should_keep):
  if not should_keep:
    return

  for line in chunk:
    sys.stdout.write(line)

for line in sys.stdin:
  if line.startswith("@"):
    sys.stdout.write(line)
    continue

  fields = line.split("\t")
  qname = fields[0]
  accession = fields[2].strip()  # RNAME column contains the accession number

  taxid = get_taxid_for_accession(accession, db_connection)

  # if debug_lookups % 1000 == 0:
  #   sys.stderr.write(
  #     f"[DEBUG] BAM Accession '{accession}' -> DB Taxid '{taxid}' (In"
  #     f" targets? {taxid in target_taxids if taxid else False})\n"
  #   )

  # debug_lookups += 1

  is_target = taxid in target_taxids if taxid else False

  if qname != current_qname:
    process_chunk(current_chunk, keep_read)
    current_qname = qname
    current_chunk = [line]
    keep_read = is_target

  else:
    current_chunk.append(line)

    if is_target:
      keep_read = True

# Process the final read chunk
process_chunk(current_chunk, keep_read)

db_connection.close()
exit()