GenerateJdbcCall.java

/*
 * SPDX-FileCopyrightText: 2026 kaumei.io
 * SPDX-License-Identifier: Apache-2.0
 */
package io.kaumei.jdbc.anno.gen;

import com.palantir.javapoet.CodeBlock;
import com.palantir.javapoet.MethodSpec;
import io.kaumei.jdbc.anno.ctx.Context;
import io.kaumei.jdbc.anno.jdbc2java.ColumnIndex;
import io.kaumei.jdbc.anno.jdbc2java.ConverterRowObjects;
import io.kaumei.jdbc.anno.jdbc2java.Jdbc2JavaConverter;
import io.kaumei.jdbc.anno.jdbc2java.Jdbc2JavaConverter.Callable;
import io.kaumei.jdbc.anno.model.JdbcTypeMirror;
import io.kaumei.jdbc.anno.model.OptionalFlag;
import io.kaumei.jdbc.anno.model.SourceMethod;
import io.kaumei.jdbc.anno.model.SourceMethodParameter;
import io.kaumei.jdbc.anno.msg.JdbcMsg;
import io.kaumei.jdbc.anno.msg.Msg;
import io.kaumei.jdbc.anno.store.SearchKey;
import io.kaumei.jdbc.anno.store.SourceDV;
import io.kaumei.jdbc.anno.utils.ParsedSql;
import io.kaumei.jdbc.core.JdbcException;
import org.jspecify.annotations.Nullable;

import javax.lang.model.element.ExecutableElement;
import java.sql.SQLException;
import java.sql.Types;
import java.util.HashSet;
import java.util.HashMap;
import java.util.Map;
import java.util.Set;

final class GenerateJdbcCall implements GenerateJdbc {

    private final Context ctx;
    private final SourceMethod.CallMethod callMethod;
    private final TargetMethod targetMethod;
    private final Msg.Builder messages;

    GenerateJdbcCall(Context ctx, SourceMethod.CallMethod sourceMethod) {
        this.ctx = ctx;
        this.callMethod = sourceMethod;
        this.targetMethod = new TargetMethod(ctx, sourceMethod.method());
        this.messages = Msg.builder();
    }

    public MethodSpec generateMethod() {
        targetMethod.beginControlFlow("try");
        targetMethod.addStatement("var con = supplier.getConnection()");
        targetMethod.addCodeBlock(ctx.kaumeiJdbcGenerator.buildSqlVariable(callMethod, targetMethod, "sql"));
        targetMethod.beginControlFlow("try (var stmt = con.prepareCall(sql))");
        var returnType = callMethod.returnType();

        var call = this.callRegister(returnType);

        this.processCallableParameter();

        targetMethod.addIfAnnotationIsPresent("stmt.setQueryTimeout($L)", callMethod.jdbcQueryTimeout());
        targetMethod.addStatement("stmt.execute()");

        if (call != null && !this.messages.hasMessages()) {
            call.converter.addCallToRow(targetMethod, "result", call.optional, call.positions);
            targetMethod.addStatement("return result");
        }

        targetMethod.endControlFlow();
        targetMethod.nextControlFlow("catch ($T e)", SQLException.class);
        targetMethod.addStatement("throw new $T(e.getMessage(), e)", JdbcException.class);
        targetMethod.endControlFlow();
        messages.add(callMethod.unusedAnno());

        return targetMethod.build(messages.build(), messages.hasMessages() ? "@JdbcCall method invalid" : "@JdbcCall method");
    }

    public void processCallableParameter() {
        var parameter = callMethod.parameter();
        var parsedSql = callMethod.parameter().parsedSql();
        for (var token : parsedSql) {
            if (token instanceof ParsedSql.SingleParameter sqlParam) {
                var name = sqlParam.root();
                var param = parameter.getOpt(name);
                if (param == null) {
                    // this is a out only parameter
                    continue;
                } else if (param instanceof SourceMethodParameter.ParamSingle sParam) {
                    var converter = ctx.kaumeiJava2Jdbc.searchJava(callMethod.method(), sParam.searchRequest());
                    if (converter.hasMessages()) {
                        messages.add(JdbcMsg.invalidSqlParameterConverter(
                                name, SourceDV.converter(ctx, param.element()), converter.messages()));
                    } else {
                        int sqlIndex = parsedSql.parameterIndexOf(sqlParam);
                        converter.value().setParameter(targetMethod, targetMethod.paramCodeBlock(name),
                                CodeBlock.of("$L", sqlIndex), sParam.optional());
                    }
                } else {
                    throw new IllegalStateException("Parameter not supported: name=" + name + ", param=" + param);
                }
            }
        }
    }

    record CallConverter(Jdbc2JavaConverter.Callable converter, OptionalFlag optional,
                         Map<String, ColumnIndex> positions) {
    }

    @Nullable CallConverter callRegister(JdbcTypeMirror<ExecutableElement> returnType) {

        return switch (returnType.kind()) {
            case VOID -> {
                validateCallablePlaceholders(Set.of());
                yield null;
            }
            case BOOLEAN, BYTE, SHORT, INT, LONG, CHAR, FLOAT, DOUBLE, RECORD, ENUM, UNKNOWN, ARRAY,
                 KAUMEI_JDBC_ROW -> {
                var request = SearchKey.of(returnType.type(), callMethod.jdbcConverterName());
                var searchResult = ctx.kaumeiJdbc2Java.searchJava(callMethod.method(), request);
                if (searchResult.hasMessages()) {
                    this.messages.add(JdbcMsg.invalidJdbcToJavaResultConverter(SourceDV.converter(ctx, callMethod.method()), request, searchResult.messages()));
                    yield null;
                }

                var converter = searchResult.value();
                if (converter instanceof Jdbc2JavaConverter.ColumnConverter c) {
                    yield callRegister_column(returnType, c);
                } else if (converter instanceof ConverterRowObjects r) {
                    yield callRegister_record(returnType, r);
                } else {
                    this.messages.add(JdbcMsg.jdbcCallOutputConverterRequiresResultSet());
                }
                yield null;

            }
            default -> {
                this.messages.add(JdbcMsg.jdbcCallRequiresSupportedReturnType(returnType));
                yield null;
            }
        };
    }

    CallConverter callRegister_record(JdbcTypeMirror<ExecutableElement> returnType,
                                      ConverterRowObjects converter) {
        var positions = new HashMap<String, ColumnIndex>();
        var parsedSql = callMethod.parameter().parsedSql();
        var outputNames = new HashSet<String>();

        converter.forAll((paramName, columnConverter) -> outputNames.add(paramName));
        validateCallablePlaceholders(outputNames);

        converter.forAll((paramName, columnConverter) -> {
            if (!parsedSql.containsName(paramName)) {
                messages.add(JdbcMsg.jdbcCallOutputPlaceholderMissing(paramName));
                return;
            }
            var tokens = parsedSql.parametersByName(paramName);
            if (tokens.size() > 1) {
                messages.add(JdbcMsg.jdbcCallOutputPlaceholderRepeated(paramName));
                return;
            }
            if (callableOutputConverter(columnConverter) != null) {
                int index = parsedSql.parameterIndexOf(tokens.get(0));
                targetMethod.addStatement("stmt.registerOutParameter($L, $L)", index, columnConverter.sqlType().sqlType());
                positions.put(paramName, ColumnIndex.ofValue(index));
            }
        });
        return new CallConverter(converter, returnType.optional(), positions);
    }


    @Nullable CallConverter callRegister_column(JdbcTypeMirror<ExecutableElement> returnType, Jdbc2JavaConverter.ColumnConverter converter) {
        var jdbcName = callMethod.jdbcName();
        if (!jdbcName.hasName()) {
            messages.add(JdbcMsg.jdbcCallScalarOutputRequiresJdbcName());
            return null;
        }
        var name = jdbcName.value();
        validateCallablePlaceholders(Set.of(name));
        if (!callMethod.parameter().parsedSql().containsName(name)) {
            messages.add(JdbcMsg.jdbcCallOutputPlaceholderMissing(name));
            return null;
        }
        var tokens = callMethod.parameter().parsedSql().parametersByName(name);
        if (tokens.size() > 1) {
            messages.add(JdbcMsg.jdbcCallOutputPlaceholderRepeated(name));
            return null;
        }
        var callable = callableOutputConverter(converter);
        if (callable == null) {
            return null;
        }
        int index = callMethod.parameter().parsedSql().parameterIndexOf(tokens.get(0));
        targetMethod.addStatement("stmt.registerOutParameter($L, $L)", index, converter.sqlType().sqlType());
        return new CallConverter(callable, returnType.optional(), Map.of(name, ColumnIndex.ofValue(index)));
    }

    private void validateCallablePlaceholders(Set<String> outputNames) {
        var parameter = callMethod.parameter();
        for (var token : parameter.parsedSql()) {
            if (token instanceof ParsedSql.UnnamedParameter) {
                messages.add(JdbcMsg.jdbcCallPlaceholderUnnamedUnsupported());
            } else if (token instanceof ParsedSql.AllValues p) {
                messages.add(JdbcMsg.jdbcCallPlaceholderExpansionUnsupported(p.name()));
            } else if (token instanceof ParsedSql.AllNames p) {
                messages.add(JdbcMsg.jdbcCallPlaceholderExpansionUnsupported(p.name()));
            } else if (token instanceof ParsedSql.Parameter p) {
                if (!p.path().isEmpty()) {
                    messages.add(JdbcMsg.jdbcCallPlaceholderExpansionUnsupported(p.name()));
                } else if (parameter.getOpt(p.root()) == null
                        && !outputNames.contains(p.root())) {
                    messages.add(JdbcMsg.jdbcCallPlaceholderNotFound(p.root()));
                }
            }
        }
    }

    private @Nullable Callable callableOutputConverter(
            Jdbc2JavaConverter.ColumnConverter converter) {
        if (!converter.hasSqlType() || !(converter instanceof Jdbc2JavaConverter.Callable callable)) {
            messages.add(JdbcMsg.jdbcCallOutputConverterRequiresResultSet());
            return null;
        }
        var sqlType = converter.sqlType().sqlType();
        if (!isSupportedCallableSqlType(sqlType)) {
            messages.add(JdbcMsg.jdbcCallOutputConverterSqlTypeUnsupported(sqlType));
            return null;
        }
        return callable;
    }

    private static boolean isSupportedCallableSqlType(int sqlType) {
        return switch (sqlType) {
            case Types.BIT, Types.BOOLEAN,
                 Types.TINYINT, Types.SMALLINT, Types.INTEGER, Types.BIGINT,
                 Types.REAL, Types.FLOAT, Types.DOUBLE,
                 Types.DECIMAL, Types.NUMERIC,
                 Types.VARCHAR, Types.CHAR, Types.LONGVARCHAR,
                 Types.NCHAR, Types.NVARCHAR, Types.LONGNVARCHAR,
                 Types.VARBINARY, Types.BINARY, Types.LONGVARBINARY,
                 Types.DATE, Types.TIME, Types.TIMESTAMP -> true;
            default -> false;
        };
    }

}