3740588efa
Refactor/daip2-9 code refactor * add missing import * refactor ac_smart_chunking.py * update tests * removed old and unused imports * removed outdated import * forgot to import re * refactor bottom up funcs * remove unused imports from file_processing.py * refactor dict_operations.py * remove unused import from conditional_funcs.py * replace top_down_funcs with one_to_n_funcs and remove top_down_funcs.py * remove duplicate import from one_to_n_funcs.py * refactor all utils.py imports * refactor error handling by removing last code in utils and integrating InvalidDateException into postprocessing_funcs * refactor import in claude_funcs.py to use string_funcs directly and (hopefully) resolve circular import * refactor regex_funcs.py to use postprocessing_funcs for add_hyphen_if_needed calls * refactor hotfix_helper_funcs.py to use postprocessing_funcs for add_hyphen_if_needed calls * refactor postprocess.py to remove unused imports * add numpy import to string_funcs * fix tests after utils refactor * move adhoc from tests to scripts * remove unused imports * remove scripts from mypy checking * isort * Merged main into refactor/daip-2-9-code-refactor Approved-by: Katon Minhas
100 lines
3.5 KiB
Python
100 lines
3.5 KiB
Python
import concurrent.futures
|
|
import traceback
|
|
import os
|
|
import csv
|
|
import time
|
|
import sys
|
|
import random
|
|
import pandas as pd
|
|
import random
|
|
random.seed(42)
|
|
|
|
import utils
|
|
import config
|
|
import claude_funcs
|
|
import file_processing
|
|
import tracking
|
|
import consolidate_output
|
|
|
|
|
|
|
|
def process_b(item):
|
|
filename, contract_text = item
|
|
if utils.contains_reimbursement(contract_text, 0):
|
|
try:
|
|
results = file_processing.run_b_prompts(item)
|
|
if results is not None:
|
|
tracking.update_results_csv(item, results)
|
|
return item, results
|
|
except Exception as e:
|
|
print(f"Error processing item {filename}: {e}")
|
|
traceback.print_exc()
|
|
else:
|
|
print(f"No reimbursement information found in {filename}. Skipping processing.")
|
|
|
|
def process_ac(item):
|
|
filename, contract_text = item
|
|
try:
|
|
results = file_processing.run_new_ac_prompts(item)
|
|
if results is not None:
|
|
tracking.update_results_csv(item, results)
|
|
return item, results
|
|
except Exception as e:
|
|
print(f"Error processing item {filename}: {e}")
|
|
traceback.print_exc()
|
|
|
|
def main():
|
|
if config.TEST:
|
|
print(claude_funcs.invoke_claude("Write 'test', nothing more.", model_id=config.MODEL_ID_CLAUDE2, filename="test", max_tokens=10))
|
|
else:
|
|
input_dict = utils.read_input() # keys are contract names, values are full contract text
|
|
total_files = len(input_dict)
|
|
|
|
print(f"Input Files : {len(input_dict)}")
|
|
|
|
# Filter no reimbursement
|
|
input_dict = {k: input_dict[k] for k in list(input_dict)[:100]}
|
|
input_dict = {k : v for k, v in input_dict.items() if utils.contains_reimbursement(str(v))} # Filter out non-contracts
|
|
print(f"Input Files after reimbursement filter: {len(input_dict)} | {total_files-len(input_dict)} files removed")
|
|
|
|
# Filter already processed
|
|
if config.FILTER_ALREADY_PROCESSED:
|
|
input_dict = utils.filter_already_processed(input_dict)
|
|
print(f"Input Files left to be processed : {len(input_dict)}")
|
|
|
|
batch_results = []
|
|
processed_count = 0
|
|
|
|
with concurrent.futures.ThreadPoolExecutor(max_workers=config.MAX_WORKERS) as executor:
|
|
if 'b' in config.FIELDS:
|
|
futures = [executor.submit(process_b, item) for item in input_dict.items()]
|
|
if 'a' in config.FIELDS and 'c' in config.FIELDS:
|
|
futures = [executor.submit(process_ac, item) for item in input_dict.items()]
|
|
|
|
for future in concurrent.futures.as_completed(futures):
|
|
try:
|
|
result = future.result()
|
|
if result:
|
|
batch_results.append(result)
|
|
processed_count += 1
|
|
|
|
if processed_count % 5 == 0:
|
|
tracking.write_batch_results(batch_results)
|
|
batch_results = []
|
|
except Exception as e:
|
|
print(f"Error in future: {e}")
|
|
traceback.print_exc()
|
|
|
|
if batch_results:
|
|
tracking.write_batch_results(batch_results)
|
|
|
|
tracking.write_stats_to_csv()
|
|
print("\nIndividual Processing Complete")
|
|
|
|
print("Consolidation starting...")
|
|
consolidate_output.consolidate_output()
|
|
print("Consolidation complete...")
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main() |