import pymc as pm
import numpy as np

def model(data):
    age = data['age']
    length = data['length']
    n_obs = len(age)
    
    with pm.Model() as model:
        # Create data container for age
        age_data = pm.Data('age', age, dims="obs_id")
        
        # Priors for von Bertalanffy growth parameters
        L_inf = pm.HalfNormal('L_inf', sigma=3.0)
        k = pm.HalfNormal('k', sigma=1.0)
        t0 = pm.Normal('t0', mu=0.0, sigma=1.0)
        
        # Expected length at each age
        mu = L_inf * (1 - pm.math.exp(-k * (age_data - t0)))
        
        # Likelihood
        sigma = pm.HalfNormal('sigma', sigma=0.5)
        y_obs = pm.Normal('y_obs', mu=mu, sigma=sigma, observed=length, dims="obs_id")
    
    return model

def gen_model(observed_data):
    built_model = model({column: observed_data[column].to_numpy() for column in observed_data.columns})
    with built_model:
        trace = pm.sample(200, tune=200, target_accept=0.90, chains=1, cores=1, random_seed=42, idata_kwargs={"log_likelihood": True})
        posterior_predictive = pm.sample_posterior_predictive(trace, random_seed=314, return_inferencedata=False)
    return built_model, posterior_predictive, trace