import socket
import threading
import ssl
from ldap3.protocol.rfc4511 import LDAPMessage, BindRequest, BindResponse, SearchRequest, SearchResultEntry, SearchResultDone
from pyasn1.codec.ber import decoder, encoder

import re
from pyasn1.codec.ber import decoder

def decode_hex_string(hex_str: str) -> str:
	"""Декодирует hex-строку (без 0x) в UTF-8 текст."""
	try:
		raw_bytes = bytes.fromhex(hex_str)
		return raw_bytes.decode('utf-8')
	except Exception:
		return None

def extract_length_from_ber(data: bytes) -> int:
	"""Извлекает ожидаемую длину ASN.1 пакета из заголовка."""
	if len(data) < 2:
		return len(data)

	length_byte = data[1]

	if length_byte < 0x80:
		return 2 + length_byte
	elif length_byte == 0x80:
		return len(data)
	else:
		num_length_bytes = length_byte & 0x7F
		if len(data) < 2 + num_length_bytes:
			return len(data)
		length = int.from_bytes(data[2:2+num_length_bytes], 'big')
		return 2 + num_length_bytes + length

def decode_ldap_full(text: str) -> str:
	"""Полное декодирование LDAP-вывода с обработкой всех форматов."""

	def decode_ber(match: re.Match) -> str:
		byte_content = match.group(1)

		# Пробуем разные уровни экранирования
		for escape_level in ['unicode_escape', 'raw']:
			try:
				if escape_level == 'unicode_escape':
					raw_bytes = byte_content.encode('utf-8').decode('unicode_escape').encode('latin-1')
				else:
					# Прямая конвертация \\xNN -> байты через regex
					def hex_replace(m):
						return chr(int(m.group(1), 16))
					raw_bytes = re.sub(r'\\x([0-9a-fA-F]{2})', hex_replace, byte_content).encode('latin-1')

				# Обрезаем по длине из ASN.1 заголовка
				expected_len = extract_length_from_ber(raw_bytes)
				if expected_len < len(raw_bytes):
					raw_bytes = raw_bytes[:expected_len]

				decoded, _ = decoder.decode(raw_bytes)
				pretty = decoded.prettyPrint()
				return f"<LDAP: {pretty}>"
			except Exception:
				continue

		# Fallback: извлекаем LDAP URL или строки
		try:
			url_match = re.search(rb'ldap://[^\x00-\x1f]+', raw_bytes)
			if url_match:
				return f"<LDAP URL: {url_match.group(0).decode('ascii')}>"
		except:
			pass

		return f"<BER parse error>"

	# Обрабатываем b'...' с любым количеством экранирований
	text = re.sub(r"b'((?:\\\\x[0-9a-fA-F]{2}|\\x[0-9a-fA-F]{2}|[^'])*)'", decode_ber, text)

	# Декодируем 0x... hex-строки в UTF-8
	def decode_hex(match: re.Match) -> str:
		hex_str = match.group(0)[2:]
		decoded = decode_hex_string(hex_str)
		if decoded:
			return f"<utf-8: {decoded}>"
		return match.group(0)

	text = re.sub(r'0x([0-9a-fA-F]{4,})', decode_hex, text)

	return text


class LDAPProxyServer:
	def __init__(self):
		self.host = '0.0.0.0'
		self.port = 3389
		self.target_host = 'practicum.cs.msu.su'
		self.target_port = 389
		self.server = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
		self.server.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
		self.server_context = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER)
		self.server_context.load_cert_chain('prac_keys/server.crt', 'prac_keys/server.key')
		self.client_context = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT)
		self.client_context.load_verify_locations('prac_keys/practicum-1.crt')
		self.client_context.check_hostname = False
		self.client_context.verify_mode = ssl.CERT_REQUIRED
		self.modified_bindpw = 'wH=8L9k!4%Ry'
		self.domen_name = 'adread@practicum.cs.msu.su'

	def start(self):
		try:
			self.server.bind((self.host, self.port))
			self.server.listen(5)
			print(f"artos_server at {self.host}:{self.port}")

			while True:
				client_socket, address = self.server.accept()
				print(f"Connected: {address}")
				ssl_client_socket = self.server_context.wrap_socket(client_socket, server_side=True)
				client_handler = threading.Thread(
					target=self.handle_client,
					args=(ssl_client_socket, address)
				)
				client_handler.start()

		except Exception as e:
			print(f"Connection error: {e}")
		finally:
			self.server.close()

	def handle_client(self, client_socket, address):
		target_socket = None
		try:
			target_socket = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
			target_socket.connect((self.target_host, self.target_port))
			print(f"Tunnel: {address} - {self.target_host}:{self.target_port}")


			def forward(source, dest, modify):
				try:
					while True:
						data = source.recv(8192)
						if not data:
							break
						if modify:
							data = self.modify(data)
						else:
							try:
								decoded, _ = decoder.decode(data, asn1Spec=LDAPMessage())
								message_type = decoded['protocolOp'].getName()
								print("<Got message type:", message_type)
								print("<Got header + rest:", decode_ldap_full(decoded.__str__()), "+", decode_ldap_full(_.__str__()))
								print("<Got end")
							except Exception as e:
								print(e)
						if data:
							dest.send(data)
				except Exception as e:
					print(f"Forwarding error: {e}")

			client_to_target = threading.Thread(
				target=forward, args=(client_socket, target_socket, True)
			)
			target_to_client = threading.Thread(
				target=forward, args=(target_socket, client_socket, False)
			)

			client_to_target.daemon = True
			target_to_client.daemon = True
			client_to_target.start()
			target_to_client.start()

			client_to_target.join()
			target_to_client.join()

		except Exception as e:
			print(f"Error: {e}")
		finally:
			if target_socket:
				target_socket.close()
			client_socket.close()
			print(f"Closed: {address}")


	def modify(self, data):
			try:
					decoded, _ = decoder.decode(data, asn1Spec=LDAPMessage())
			except Exception as e:
					return data
			message_type = decoded['protocolOp'].getName()
			print(">Sended message type:", message_type)
			print(">Sended header + rest:", decode_ldap_full(decoded.__str__()), "+", decode_ldap_full(_.__str__()))
			print(">Sended end")
			match message_type:
					case 'bindRequest':
							if decoded['protocolOp']['bindRequest']['name'] == self.domen_name.encode():
									decoded['protocolOp']['bindRequest']['authentication']['simple'] = self.modified_bindpw.encode()
					case 'searchRequest':
							search = decoded['protocolOp']['searchRequest']
							ban = True
							# print([search])
							print(str(search['baseObject']))
							print(str(search['scope']))
							print(str(search['filter']['present']))
							print([str(i) for i in search['attributes']])
							# if str(search['baseObject']) == 'dc=PRACTICUM,dc=CS,dc=MSU,dc=SU':
							if str(search['scope']) == 'baseObject':
									if str(search['filter']['present']) == 'objectClass':
											if [str(attr) for attr in search['attributes']] == ['dn']:
													ban = False
							if ban:
									print("Banned search message:", decoded, _)
									return False
					case 'unbindRequest' | 'abandonRequest':
							pass
					case _:
							print("Banned message type:", _)
							return False
			modified_data = encoder.encode(decoded)
			return modified_data

if __name__ == "__main__":
	server = LDAPProxyServer()
	server.start()

