# One kernel: a namespace, and a loop that runs code in it.
#
# Requests arrive as one JSON object per line on stdin, answers leave as one JSON object per line
# on stdout. Nothing else is ever written to the real stdout — the cell's own prints are captured
# and travel inside the answer — so the two ends never have to tell a frame from a print.
#
# The last value is the answer and there is no `print` to write, which is the rule the JavaScript
# cell already has: every statement but the last is exec'd, and the last is eval'd when it is an
# expression.

import ast
import base64
import io
import json
import os
import signal
import sys
import traceback

# how many rows a result may carry back. The whole answer crosses a socket as JSON, and this is
# the same ceiling the duckdb router puts on a query.
MAX_ROWS = 10000

namespace = {"__name__": "__main__"}


# pandas: the whole frame is already in memory, and to_dict("records") is all of it.
def is_frame(value):
	return type(value).__name__ == "DataFrame" and hasattr(value, "to_dict")


# polars: also a DataFrame, also in memory, and it answers to to_dict as well — but with a different
# signature, so it has to be told apart before pandas is asked. to_dicts, plural, is the one only it
# has, and is the call that gives the rows.
def is_polars_frame(value):
	return type(value).__name__ == "DataFrame" and hasattr(value, "to_dicts")


# spark: a DataFrame is a plan rather than rows, and it may stand for more of them than this machine
# holds. Told apart by what only it has — pandas has neither collect nor limit, and polars has limit
# without collect.
def is_spark_frame(value):
	return type(value).__name__ == "DataFrame" and hasattr(value, "collect") and hasattr(value, "limit")


def is_series(value):
	return type(value).__name__ == "Series" and hasattr(value, "tolist")


def is_figure(value):
	return type(value).__name__ == "Figure" and hasattr(value, "savefig")


def is_rows(value):
	return isinstance(value, list) and len(value) > 0 and isinstance(value[0], dict)


# json cannot carry a numpy int, a Timestamp or a NaN, and a cell that returns one of those is
# entirely ordinary. Everything that is not already a json value becomes its str.
def plain(value):
	if value is None or isinstance(value, (bool, int, float, str)):
		if isinstance(value, float) and (value != value or value in (float("inf"), float("-inf"))):
			return None
		else:
			return value
	else:
		return str(value)


def rows_of(records):
	rows = []
	for record in records[:MAX_ROWS]:
		rows.append({str(key): plain(value) for key, value in record.items()})
	return rows


def png_of(figure):
	buffer = io.BytesIO()
	figure.savefig(buffer, format="png", bbox_inches="tight")
	return base64.b64encode(buffer.getvalue()).decode("ascii")


# What the cell drew, when it drew without returning the figure. pyplot keeps the figures it made,
# and a cell that ends in plt.plot(...) has one open with nothing pointing at it.
def open_figure():
	pyplot = sys.modules.get("matplotlib.pyplot")
	if pyplot is not None and pyplot.get_fignums():
		figure = pyplot.figure(pyplot.get_fignums()[-1])
		png = png_of(figure)
		pyplot.close("all")
		return png
	else:
		return None


# One page of a spark frame, and never more. limit() is part of the plan, so the job that runs reads
# only what is asked for — collect() on the frame itself would pull every row of a table that has no
# reason to fit here. One row past the ceiling is fetched so "there is more" can be said without
# counting the rest, which would be a second job over the whole thing.
def spark_result(frame):
	taken = frame.limit(MAX_ROWS + 1).collect()
	records = [row.asDict(recursive=True) for row in taken[:MAX_ROWS]]
	return {"kind": "rows", "rows": rows_of(records), "truncated": len(taken) > MAX_ROWS}


def result_of(value):
	# polars before pandas: both are called DataFrame and both have a to_dict, so the more specific
	# test has to run first or a polars frame is read with pandas's call and raises a TypeError.
	if is_polars_frame(value):
		return {"kind": "rows", "rows": rows_of(value.to_dicts()), "truncated": len(value) > MAX_ROWS}
	elif is_frame(value):
		return {"kind": "rows", "rows": rows_of(value.to_dict("records")), "truncated": len(value) > MAX_ROWS}
	elif is_spark_frame(value):
		return spark_result(value)
	elif is_series(value):
		listed = value.tolist()[:MAX_ROWS]
		return {"kind": "rows", "rows": [{"index": plain(key), "value": plain(item)} for key, item in enumerate(listed)], "truncated": False}
	elif is_figure(value):
		return {"kind": "image", "png": png_of(value)}
	elif is_rows(value):
		return {"kind": "rows", "rows": rows_of(value), "truncated": len(value) > MAX_ROWS}
	elif value is None:
		drawn = open_figure()
		if drawn is None:
			return {"kind": "none"}
		else:
			return {"kind": "image", "png": drawn}
	else:
		return {"kind": "text", "text": repr(value)}


# Every statement but the last is exec'd; the last is eval'd when it is an expression. That is what
# makes a cell of one line and a cell of ten the same thing, with the last line as its answer.
def run(code):
	parsed = ast.parse(code, "<cell>", "exec")
	if len(parsed.body) > 0 and isinstance(parsed.body[-1], ast.Expr):
		head = ast.Module(body=parsed.body[:-1], type_ignores=[])
		tail = ast.Expression(body=parsed.body[-1].value)
		exec(compile(head, "<cell>", "exec"), namespace)
		return eval(compile(tail, "<cell>", "eval"), namespace)
	else:
		exec(compile(parsed, "<cell>", "exec"), namespace)
		return None


# A SparkSession whose JVM is gone is a session getOrCreate() cannot rebuild. pyspark caches the
# context and the gateway on the class, hands the dead one back, and every call on it fails on a
# socket that nothing is listening to — so the kernel stays broken until it is restarted, and there
# is no way to ask for a session again.
#
# The JVM is a child process of this one, so whether it is still there is a poll() rather than a
# round trip. When it has exited, the cached handles are dropped and the next getOrCreate() launches
# a new JVM, which is what asking for a session meant in the first place.
def clear_dead_spark():
	context = getattr(sys.modules.get("pyspark"), "SparkContext", None)
	gateway = getattr(context, "_gateway", None)
	process = getattr(gateway, "proc", None)
	if process is not None and process.poll() is not None:
		context._gateway = None
		context._jvm = None
		context._active_spark_context = None
		session = sys.modules.get("pyspark.sql")
		if session is not None:
			session.SparkSession._instantiatedSession = None
			session.SparkSession._activeSession = None
		else:
			# pyspark.sql was never imported, so there is no session to forget
			pass
	else:
		# pyspark is not here, there is no JVM yet, or the one there is is still running
		pass


# The JVM this kernel started, while it is still running. It is a child process of this one, which
# is what lets the app kill it if it outlives the kernel — reported with every answer so the pid is
# known on the other side before it is needed.
def spark_jvm_pid():
	context = getattr(sys.modules.get("pyspark"), "SparkContext", None)
	process = getattr(getattr(context, "_gateway", None), "proc", None)
	if process is not None and process.poll() is None:
		return process.pid
	else:
		return None


# SIGTERM is how this kernel is stopped — a venv change, a venv deleted, the app closing. A
# SparkSession left behind holds a JVM with gigabytes in it and nothing else would ever tell it to
# stop, so it is stopped here rather than left to a socket teardown to notice.
#
# The alarm is the bound on that: stop() talks to the JVM over a socket, and a JVM that has stopped
# answering would otherwise hold this handler open for as long as it liked.
def stop_spark_and_exit(number, frame):
	signal.alarm(5)
	context = getattr(sys.modules.get("pyspark"), "SparkContext", None)
	active = getattr(context, "_active_spark_context", None)
	if active is not None:
		try:
			active.stop()
		except BaseException:
			# already gone, or not answering. The app kills the JVM by pid either way.
			pass
	else:
		# no session in this kernel, so there is nothing holding a JVM
		pass
	os._exit(0)


def execute(request):
	clear_dead_spark()
	captured = io.StringIO()
	real_stdout = sys.stdout
	real_stderr = sys.stderr
	sys.stdout = captured
	sys.stderr = captured
	answer = {}
	try:
		if "oe" in namespace:
			namespace["oe"].current = request.get("object") or {}
		value = run(request.get("code") or "")
		answer = {"ok": True, "result": result_of(value)}
	except BaseException:
		answer = {"ok": False, "error": traceback.format_exc()}
	finally:
		sys.stdout = real_stdout
		sys.stderr = real_stderr
	answer["stdout"] = captured.getvalue()
	answer["jvm"] = spark_jvm_pid()
	return answer


# A one-line summary of a value, so the variables list says what a name holds rather than only that
# it holds something. A DataFrame says its shape, which is the one thing worth reading about one.
def summary_of(value):
	if is_frame(value) or is_polars_frame(value):
		return "%d rows x %d columns" % value.shape
	elif is_spark_frame(value):
		# the columns and not the rows: a schema is already known, and counting rows means running
		# the whole query to write one line in a list nobody asked to wait for
		return "%d columns" % len(value.columns)
	elif is_series(value):
		return "%d values" % len(value)
	elif isinstance(value, (list, tuple, set, dict)):
		return "%d items" % len(value)
	else:
		text = repr(value)
		if len(text) > 80:
			return text[:80] + "…"
		else:
			return text


# The contents, not the furniture: dunder names, imported modules and oe itself are what every
# kernel has and nothing anyone put there.
def variables():
	listed = []
	for name, value in namespace.items():
		if name.startswith("__") or name == "oe" or type(value).__name__ == "module":
			continue
		listed.append({"name": name, "type": type(value).__name__, "summary": summary_of(value)})
	listed.sort(key=lambda entry: entry["name"])
	return listed


def answer(payload):
	sys.__stdout__.write(json.dumps(payload) + "\n")
	sys.__stdout__.flush()


def main():
	try:
		import oe

		namespace["oe"] = oe
	except BaseException:
		# a kernel with no oe still runs code; the cell says so when a call to it fails
		pass
	signal.signal(signal.SIGTERM, stop_spark_and_exit)
	answer({"op": "ready"})
	for line in sys.stdin:
		text = line.strip()
		if text:
			request = json.loads(text)
			op = request.get("op")
			if op == "execute":
				answer({"id": request.get("id"), **execute(request)})
			elif op == "variables":
				answer({"id": request.get("id"), "ok": True, "variables": variables()})
			else:
				answer({"id": request.get("id"), "ok": True})
		else:
			pass


main()
