diff --git a/app.py b/app.py index 16afbe6..047db12 100644 --- a/app.py +++ b/app.py @@ -111,19 +111,53 @@ def get_file_key(filepath): return filepath def is_file_protected(filepath): - """Check if a file is password protected.""" + """Check if a file is password protected. Checks both encrypted and non-encrypted versions.""" passwords = load_passwords() file_key = get_file_key(filepath) - return passwords.get(file_key, {}).get('protected', False) + + # Check if this exact key is protected + if passwords.get(file_key, {}).get('protected', False): + return True + + # If looking for encrypted version, also check non-encrypted key + if filepath.endswith('.enc'): + base_key = filepath[:-4] # Remove .enc + if passwords.get(base_key, {}).get('protected', False): + return True + + # If looking for non-encrypted, also check encrypted key + else: + enc_key = filepath + '.enc' + if passwords.get(enc_key, {}).get('protected', False): + return True + + return False def is_authenticated(filepath): """Check if user is authenticated for a protected file.""" if not is_file_protected(filepath): return True # Not protected, so allowed + passwords = load_passwords() file_key = get_file_key(filepath) authenticated_files = session.get('authenticated_files', {}) - return authenticated_files.get(file_key, False) + + # Check if authenticated for this exact key + if authenticated_files.get(file_key, False): + return True + + # If looking for encrypted version, also check non-encrypted key + if filepath.endswith('.enc'): + base_key = filepath[:-4] + if authenticated_files.get(base_key, False): + return True + # If looking for non-encrypted, also check encrypted key + else: + enc_key = filepath + '.enc' + if authenticated_files.get(enc_key, False): + return True + + return False def extract_email_text(email_content): """ @@ -631,11 +665,16 @@ def view_file(filepath): with open(file_path, 'r', encoding='utf-8') as f: encrypted_content = f.read() - # Get password from session + # Get password from session - check multiple possible keys file_key = get_file_key(lookup_filepath) authenticated_files = session.get('authenticated_files', {}) password = authenticated_files.get(file_key + '_password') + # If not found with lookup_filepath key, try base filename (without .enc) + if not password and lookup_filepath.endswith('.enc'): + base_key = lookup_filepath[:-4] + password = authenticated_files.get(base_key + '_password') + if not password: return "Unable to decrypt: password not found in session", 500 @@ -685,19 +724,34 @@ def authenticate(filepath): passwords = load_passwords() file_data = passwords.get(file_key) + actual_key = file_key + + # If looking for .enc file, also check the base name + if not file_data and filepath.endswith('.enc'): + base_key = filepath[:-4] + file_data = passwords.get(base_key) + if file_data: + actual_key = base_key + + # If looking for base file, also check the .enc version + if not file_data and not filepath.endswith('.enc'): + enc_key = filepath + '.enc' + file_data = passwords.get(enc_key) + if file_data: + actual_key = enc_key if not file_data or not file_data.get('protected'): return jsonify({'success': False, 'error': 'File not protected'}), 400 # Check password if check_password_hash(file_data['password_hash'], password): - # Store in session + # Store in session using the key where protection is stored if 'authenticated_files' not in session: session['authenticated_files'] = {} - session['authenticated_files'][file_key] = True + session['authenticated_files'][actual_key] = True # Store the password for decryption if file is encrypted if is_encrypted_file(filepath): - session['authenticated_files'][file_key + '_password'] = password + session['authenticated_files'][actual_key + '_password'] = password session.modified = True logger.info(f"User authenticated for {file_key}")